diff --git a/.clang-format b/.clang-format index 1c7255f3..e793082f 100644 --- a/.clang-format +++ b/.clang-format @@ -53,6 +53,7 @@ IndentWrappedFunctionNames: false KeepEmptyLinesAtTheStartOfBlocks: false MaxEmptyLinesToKeep: 1 NamespaceIndentation: None +PenaltyBreakBeforeMemberAccess: 1000 PenaltyReturnTypeOnItsOwnLine: 400 PointerAlignment: Left ReflowComments: false diff --git a/.github/workflows/cmake-cuda-multi-compiler.yml b/.github/workflows/cmake-cuda-multi-compiler.yml index 8365d948..872cbd0b 100644 --- a/.github/workflows/cmake-cuda-multi-compiler.yml +++ b/.github/workflows/cmake-cuda-multi-compiler.yml @@ -13,11 +13,11 @@ jobs: runs-on: ${{ matrix.os }} strategy: - fail-fast: true + fail-fast: false matrix: os: [ubuntu-latest] - build_type: [Debug] + build_type: [Release] cpp_compiler: [g++, clang++] steps: @@ -50,8 +50,9 @@ jobs: -DKMM_BUILD_TESTS:BOOL=ON -DKMM_BUILD_EXAMPLES:BOOL=ON -DKMM_BUILD_BENCHMARKS:BOOL=ON + -DKMM_ENABLE_WARNINGS:BOOL=OFF -S ${{ github.workspace }} - name: Build # Build your program with the given configuration. Note that --config is needed because the default Windows generator is a multi-config generator (Visual Studio generator). - run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} --parallel 4 + run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} diff --git a/.github/workflows/cmake-hip.yml b/.github/workflows/cmake-hip.yml index b33981be..65127112 100644 --- a/.github/workflows/cmake-hip.yml +++ b/.github/workflows/cmake-hip.yml @@ -13,22 +13,27 @@ jobs: runs-on: ${{ matrix.os }} strategy: - fail-fast: true + fail-fast: false matrix: os: [ubuntu-latest] - build_type: [Debug] + build_type: [Release] steps: - - uses: loostrum/rocm-installer@main - with: - version: 6.2.2 - packages: hip-dev hipcc rocm-device-libs rocblas rocm-hip-runtime-dev rocrand - - uses: actions/checkout@v4 with: submodules: 'recursive' + - name: Install dependencies + run: | + sudo apt update + sudo apt install python3-setuptools python3-wheel + wget https://repo.radeon.com/amdgpu-install/6.3.4/ubuntu/noble/amdgpu-install_6.3.60304-1_all.deb + sudo apt install ./amdgpu-install_6.3.60304-1_all.deb + rm ./amdgpu-install_6.3.60304-1_all.deb + sudo apt update + sudo apt install hip-dev hipcc rocm-device-libs rocblas rocm-hip-runtime-dev rocrand hipcub-dev + - name: Set reusable strings # Turn repeated input strings (such as the build output directory) into step outputs. These step outputs can be used throughout the workflow file. id: strings @@ -42,9 +47,11 @@ jobs: run: > cmake -B ${{ steps.strings.outputs.build-output-dir }} -DCMAKE_PREFIX_PATH=/opt/rocm + -DCMAKE_CXX_COMPILER:PATH=/opt/rocm/bin/amdclang++ -DCMAKE_HIP_COMPILER_ROCM_ROOT:PATH=/opt/rocm -DCMAKE_HIP_ARCHITECTURES=gfx90a -DAMDGPU_TARGETS=gfx90a + -DCMAKE_HIP_COMPILER:PATH=/opt/rocm/bin/amdclang++ -DCMAKE_BUILD_TYPE=${{ matrix.build_type }} -DCMAKE_VERBOSE_MAKEFILE:BOOL=ON -DSPDLOG_FMT_EXTERNAL:BOOL=ON @@ -52,8 +59,9 @@ jobs: -DKMM_BUILD_TESTS:BOOL=ON -DKMM_BUILD_EXAMPLES:BOOL=ON -DKMM_BUILD_BENCHMARKS:BOOL=ON + -DKMM_ENABLE_WARNINGS:BOOL=OFF -S ${{ github.workspace }} - name: Build # Build your program with the given configuration. Note that --config is needed because the default Windows generator is a multi-config generator (Visual Studio generator). - run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} --parallel 4 + run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} diff --git a/.github/workflows/cmake-multi-compiler.yml b/.github/workflows/cmake-multi-compiler.yml index 1063d0ff..0a5bbba5 100644 --- a/.github/workflows/cmake-multi-compiler.yml +++ b/.github/workflows/cmake-multi-compiler.yml @@ -13,11 +13,11 @@ jobs: runs-on: ${{ matrix.os }} strategy: - fail-fast: true + fail-fast: false matrix: os: [ubuntu-latest, macos-latest] - build_type: [Debug] + build_type: [Release] cpp_compiler: [g++, clang++] steps: @@ -42,11 +42,12 @@ jobs: -DCMAKE_VERBOSE_MAKEFILE:BOOL=ON -DSPDLOG_FMT_EXTERNAL:BOOL=ON -DKMM_BUILD_TESTS:BOOL=ON + -DKMM_ENABLE_WARNINGS:BOOL=OFF -S ${{ github.workspace }} - name: Build # Build your program with the given configuration. Note that --config is needed because the default Windows generator is a multi-config generator (Visual Studio generator). - run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} --parallel 4 + run: cmake --build ${{ steps.strings.outputs.build-output-dir }} --config ${{ matrix.build_type }} - name: Test working-directory: ${{ steps.strings.outputs.build-output-dir }} diff --git a/.gitignore b/.gitignore index a8500b20..a1e172af 100644 --- a/.gitignore +++ b/.gitignore @@ -5,3 +5,6 @@ cmake-build* # Editors .idea/ .vscode/ + +# Generated docs +docs/html/ diff --git a/.gitmodules b/.gitmodules index 01d96300..8040d975 100644 --- a/.gitmodules +++ b/.gitmodules @@ -7,3 +7,6 @@ [submodule "external/Catch2"] path = external/Catch2 url = https://github.com/catchorg/Catch2.git +[submodule "external/unordered_dense"] + path = external/unordered_dense + url = https://github.com/martinus/unordered_dense diff --git a/CMakeLists.txt b/CMakeLists.txt index 5a930d1c..6b14a3b3 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -2,7 +2,7 @@ cmake_minimum_required(VERSION 3.10) # Project setup set(PROJECT_NAME "kmm") -project(${PROJECT_NAME} LANGUAGES CXX VERSION 0.3) +project(${PROJECT_NAME} LANGUAGES CXX VERSION 0.4.0) # User options (features) option(KMM_STATIC "Build a static library" OFF) @@ -14,6 +14,7 @@ option(KMM_ENABLE_LINTER "Enable clang-tidy linter" OFF) option(KMM_BUILD_TESTS "Build tests" OFF) option(KMM_BUILD_EXAMPLES "Build examples" OFF) option(KMM_BUILD_BENCHMARKS "Build benchmarks" OFF) +option(KMM_ENABLE_WARNINGS "Enable -Wall -Wextra -Wconversion" ON) # Check options if(KMM_USE_CUDA AND KMM_USE_HIP) @@ -29,8 +30,11 @@ file(GLOB_RECURSE sources "${PROJECT_SOURCE_DIR}/src/*.cu" ) -# Exclude backend implementation files from compilation, as they are included in `src/core/backends.cpp` -list(FILTER sources EXCLUDE REGEX "src/backends/(cuda|hip|cpu)\\.cpp$") +list(FILTER sources EXCLUDE REGEX "src/utils/backends/cpu.cpp") + +if(NOT KMM_USE_CUDA AND NOT KMM_USE_HIP) + list(APPEND sources "${PROJECT_SOURCE_DIR}/src/utils/backends/cpu.cpp") +endif() # Create library if(KMM_STATIC) @@ -53,7 +57,8 @@ endif() target_compile_options(${PROJECT_NAME} PUBLIC $<$:-forward-unknown-to-host-compiler> - -Wall -Wextra -Wconversion -Wno-unused-parameter + "$<$:-Wall;-Wextra;-Wconversion;-Wno-unused-parameter>" + $<$>:-w> #$<$:-Xcompiler=-Werror> ) target_compile_options(${PROJECT_NAME} PUBLIC ${CXXFLAGS}) @@ -73,6 +78,10 @@ set(SPDLOG_BUILD_PIC ON) add_subdirectory(external/spdlog) target_link_libraries(${PROJECT_NAME} PUBLIC spdlog) +# Dependencies: unordered_dense +add_subdirectory(external/unordered_dense) +target_link_libraries(${PROJECT_NAME} PUBLIC unordered_dense::unordered_dense) + # Install include(GNUInstallDirs) set_target_properties( @@ -143,13 +152,31 @@ elseif(KMM_USE_HIP) set(CMAKE_MODULE_PATH "${HIP_PATH}/cmake" ${CMAKE_MODULE_PATH}) find_package(HIP REQUIRED) find_package(ROCBLAS REQUIRED) - set_source_files_properties(${sources} PROPERTIES LANGUAGE HIP) + find_package(hipcub REQUIRED) + # CMake only recognizes .cu as CUDA by default, so the actual GPU kernel files (which need to be + # compiled by amdclang++'s HIP frontend) must be reassigned explicitly. Plain .cpp files must + # stay as CXX -- forcing them to HIP too makes amdclang++ run a device-code compilation pass on + # files with no device code, which fails since they never include . + set(hip_kernel_sources ${sources}) + list(FILTER hip_kernel_sources INCLUDE REGEX "\\.cu$") + set_source_files_properties(${hip_kernel_sources} PROPERTIES LANGUAGE HIP) + + # `enable_language(HIP)` above means CMake's native HIP-language support already knows how to + # compile the HIP-language kernel sources (invoking amdclang++ with the right -x hip flag and + # architecture), the same way CUDA_ARCHITECTURES drives .cu compilation in the CUDA branch. We + # deliberately do NOT link hip::device here: it exists for the legacy pattern of writing device + # code directly in .cpp files, and does so by force-injecting -x hip onto every CXX-language + # source in the target -- which breaks plain host .cpp files that never include hip headers. + set_target_properties( + ${PROJECT_NAME} + PROPERTIES + HIP_ARCHITECTURES "gfx90a" + ) target_link_libraries(${PROJECT_NAME} PUBLIC hip::host roc::rocblas - PRIVATE - hip::device + hip::hipcub ) # Define `KMM_USE_HIP` macro so that headers can detect HIP usage @@ -171,4 +198,4 @@ endif() # Compile benchmarks if(KMM_BUILD_BENCHMARKS) add_subdirectory(benchmarks) -endif() +endif() \ No newline at end of file diff --git a/LICENSE b/LICENSE index 261eeb9e..57bc88a1 100644 --- a/LICENSE +++ b/LICENSE @@ -199,3 +199,4 @@ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions and limitations under the License. + diff --git a/Makefile b/Makefile index f1fa8fd0..8794c12a 100644 --- a/Makefile +++ b/Makefile @@ -1,12 +1,13 @@ CLANG_FORMAT=clang-format --verbose pretty: - ${CLANG_FORMAT} -i include/kmm/*.hpp include/kmm/*/*.hpp - $(CLANG_FORMAT) -i src/*/*.cpp src/*/*.cu src/*/*.cuh - ${CLANG_FORMAT} -i test/*/*.cpp - ${CLANG_FORMAT} -i examples/*.cu - ${CLANG_FORMAT} -i benchmarks/*.cu + find include src test examples \ + \( -name '*.hpp' -o -name '*.cpp' -o -name '*.cu' -o -name '*.cuh' \) -type f -print0 \ + | xargs -0 -r ${CLANG_FORMAT} -i + +docs: + cd docs && doxygen Doxyfile all: pretty -.PHONY : pretty +.PHONY : pretty docs diff --git a/README.md b/README.md index 51de197b..e17ecf2a 100644 --- a/README.md +++ b/README.md @@ -84,4 +84,5 @@ int main() { ## License -KMM is made available under the terms of the Apache License version 2.0, see the file LICENSE for details. +KMM is made available under the terms of the Apache License version 2.0, see the LICENSE file for details. + diff --git a/benchmarks/vector_add.cu b/benchmarks/vector_add.cu index 8b69ab38..1713859f 100644 --- a/benchmarks/vector_add.cu +++ b/benchmarks/vector_add.cu @@ -1,29 +1,25 @@ #include #include #include +#include #include "kmm/kmm.hpp" -#include "kmm/utils/integer_fun.hpp" using real_type = float; const unsigned int max_iterations = 10; -__global__ void initialize_range(kmm::Range chunk, kmm::GPUSubviewMut output) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + chunk.begin; - if (i >= chunk.end) { +__global__ void initialize_range(int64_t offset, kmm::ViewMut output) { + int64_t i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= output.size()) { return; } - output[i] = static_cast(i); + output[i] = static_cast(offset + i); } -__global__ void fill_range( - kmm::Range chunk, - real_type value, - kmm::GPUSubviewMut output -) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + chunk.begin; - if (i >= chunk.end) { +__global__ void fill_range(real_type value, kmm::ViewMut output) { + int64_t i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= output.size()) { return; } @@ -31,88 +27,115 @@ __global__ void fill_range( } __global__ void vector_add( - kmm::Range range, - kmm::GPUSubviewMut output, - kmm::GPUSubview left, - kmm::GPUSubview right + kmm::ViewMut output, + kmm::View left, + kmm::View right ) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; + int64_t i = blockIdx.x * blockDim.x + threadIdx.x; - if (i >= range.end) { + if (i >= output.size()) { return; } output[i] = left[i] + right[i]; } +// Launches `kernel` (previously bound to a fixed block size) over `extent` elements, computing +// the grid size to cover it. +template +void launch( + kmm::Device& device, + K kernel, + unsigned int block_size, + int64_t extent, + Args&&... args +) { + auto grid_size = static_cast(kmm::div_ceil(extent, int64_t(block_size))); + + device.submit(kmm::Kernel(kernel, dim3(grid_size), dim3(block_size)), std::forward(args)...); +} + bool inner_loop( - kmm::RuntimeHandle& rt, + kmm::Context& context, + std::vector& devices, unsigned int threads, int64_t n, int64_t chunk_size, std::chrono::duration& init_time, std::chrono::duration& run_time ) { - using namespace kmm::placeholders; - - dim3 block_size = threads; auto timing_start_init = std::chrono::steady_clock::now(); - auto A = kmm::Array {n}; - auto B = kmm::Array {n}; - auto C = kmm::Array {n}; + + // Each chunk of `A`/`B`/`C` is its own array (own buffer, own home device), which is what + // lets `num_chunks` independent submissions below actually run concurrently/pipelined, + // rather than serializing on one shared buffer. + kmm::Distribution<1> dist{kmm::Shape<1>(n), kmm::Shape<1>(chunk_size)}; + kmm::DistArray A(context.runtime(), dist); + kmm::DistArray B(context.runtime(), dist); + kmm::DistArray C(context.runtime(), dist); // Initialize input arrays - rt.parallel_submit( - kmm::TileDomain(n, chunk_size), - kmm::GPUKernel(initialize_range, block_size), - _x, - write(A[_x]) - ); + for (size_t i = 0; i < dist.num_chunks(); i++) { + auto& device = devices[A.chunk_home(i).as_device().get()]; + int64_t offset = dist.chunk_offset(dist.unravel(i))[0]; + int64_t extent = A.chunk(i).size(); - rt.parallel_submit( - kmm::TileDomain(n, chunk_size), - kmm::GPUKernel(fill_range, block_size), - _x, - static_cast(1.0), - write(B[_x]) - ); + launch(device, initialize_range, threads, extent, offset, kmm::write(A.chunk(i))); + launch(device, fill_range, threads, extent, static_cast(1.0), kmm::write(B.chunk(i))); + } - rt.synchronize(); + context.synchronize(); auto timing_stop_init = std::chrono::steady_clock::now(); init_time += timing_stop_init - timing_start_init; // Benchmark auto timing_start = std::chrono::steady_clock::now(); - rt.parallel_submit( - kmm::TileDomain(n, chunk_size), - kmm::GPUKernel(vector_add, block_size), - _x, - write(C[_x]), - A[_x], - B[_x] - ); - rt.synchronize(); + for (size_t i = 0; i < dist.num_chunks(); i++) { + auto& device = devices[C.chunk_home(i).as_device().get()]; + int64_t extent = C.chunk(i).size(); + + launch(device, vector_add, threads, extent, kmm::write(C.chunk(i)), A.chunk(i), B.chunk(i)); + } + + context.synchronize(); auto timing_stop = std::chrono::steady_clock::now(); run_time += timing_stop - timing_start; // Correctness check - std::vector result(n); - C.copy_to(result); - - for (unsigned int i = 0; i < n; i++) { - if (result[i] != static_cast(i) + 1) { - std::cerr << "Wrong result at " << i << " : " << result[i] - << " != " << static_cast(i) + 1.0 << std::endl; - return false; + bool status = true; + + for (size_t i = 0; i < dist.num_chunks(); i++) { + int64_t offset = dist.chunk_offset(dist.unravel(i))[0]; + std::vector result = context.to_vector(C.chunk(i)); + + for (size_t j = 0; j < result.size(); j++) { + int64_t global_i = offset + static_cast(j); + auto expected = static_cast(global_i) + static_cast(1.0); + + if (result[j] != expected) { + std::cerr << "Wrong result at " << global_i << " : " << result[j] + << " != " << expected << std::endl; + status = false; + } } } - return true; + return status; } int main(int argc, char* argv[]) { - auto rt = kmm::make_runtime(); + auto config = kmm::default_config_from_environment(); + kmm::Context context = kmm::make_runtime(config); + + // One `Device` per physical GPU, created once and reused across chunks/iterations -- + // constructing a `Device` owns a GPU stream, so it should not happen inside the hot loop. + size_t num_devices = context.system_info().num_devices(); + std::vector devices; + for (size_t i = 0; i < num_devices; i++) { + devices.push_back(context.gpu(kmm::DeviceId(i))); + } + bool status = false; int64_t n = 0; int64_t num_chunks = 0; @@ -133,8 +156,15 @@ int main(int argc, char* argv[]) { mem *= double(n); // Warm-up run - status = - inner_loop(rt, num_threads, n, kmm::div_ceil(n, num_chunks), init_time, vector_add_time); + status = inner_loop( + context, + devices, + num_threads, + n, + kmm::div_ceil(n, num_chunks), + init_time, + vector_add_time + ); if (!status) { std::cerr << "Warm-up run failed." << std::endl; return 1; @@ -145,7 +175,8 @@ int main(int argc, char* argv[]) { for (unsigned int iteration = 0; iteration < max_iterations; ++iteration) { status = inner_loop( - rt, + context, + devices, num_threads, n, kmm::div_ceil(n, num_chunks), diff --git a/examples/stencil.cu b/docs/.gitkeep similarity index 100% rename from examples/stencil.cu rename to docs/.gitkeep diff --git a/docs/Doxyfile b/docs/Doxyfile index 7cd416fe..9aced25a 100644 --- a/docs/Doxyfile +++ b/docs/Doxyfile @@ -1,4 +1,4 @@ -# Doxyfile 1.10.0 +# Doxyfile 1.9.1 # This file describes the settings to be used by the documentation system # doxygen (www.doxygen.org) for a project. @@ -12,16 +12,6 @@ # For lists, items can also be appended using: # TAG += value [value, ...] # Values that contain spaces should be placed between quotes (\" \"). -# -# Note: -# -# Use doxygen to compare the used configuration file with the template -# configuration file: -# doxygen -x [configFile] -# Use doxygen to compare the used configuration file with the template -# configuration file without replacing the environment variables or CMake type -# replacement variables: -# doxygen -x_noenv [configFile] #--------------------------------------------------------------------------- # Project related configuration options @@ -42,19 +32,19 @@ DOXYFILE_ENCODING = UTF-8 # title of most generated pages and in a few other places. # The default value is: My Project. -PROJECT_NAME = "KMM" +PROJECT_NAME = "kmm" # The PROJECT_NUMBER tag can be used to enter a project or revision number. This # could be handy for archiving the generated documentation or if some version # control system is used. -PROJECT_NUMBER = +PROJECT_NUMBER = 0.1.0 # Using the PROJECT_BRIEF tag one can provide an optional one line description # for a project that appears at the top of each page and should give viewer a # quick idea about the purpose of the project. Keep the description short. -PROJECT_BRIEF = +PROJECT_BRIEF = "Kernel launcher framework for heterogeneous systems" # With the PROJECT_LOGO tag one can specify a logo or an icon that is included # in the documentation. The maximum height of the logo should not exceed 55 @@ -63,41 +53,23 @@ PROJECT_BRIEF = PROJECT_LOGO = -# With the PROJECT_ICON tag one can specify an icon that is included in the tabs -# when the HTML document is shown. Doxygen will copy the logo to the output -# directory. - -PROJECT_ICON = - # The OUTPUT_DIRECTORY tag is used to specify the (relative or absolute) path # into which the generated documentation will be written. If a relative path is # entered, it will be relative to the location where doxygen was started. If # left blank the current directory will be used. -OUTPUT_DIRECTORY = _doxygen +OUTPUT_DIRECTORY = -# If the CREATE_SUBDIRS tag is set to YES then doxygen will create up to 4096 -# sub-directories (in 2 levels) under the output directory of each output format -# and will distribute the generated files over these directories. Enabling this +# If the CREATE_SUBDIRS tag is set to YES then doxygen will create 4096 sub- +# directories (in 2 levels) under the output directory of each output format and +# will distribute the generated files over these directories. Enabling this # option can be useful when feeding doxygen a huge amount of source files, where # putting all generated files in the same directory would otherwise causes -# performance problems for the file system. Adapt CREATE_SUBDIRS_LEVEL to -# control the number of sub-directories. +# performance problems for the file system. # The default value is: NO. CREATE_SUBDIRS = NO -# Controls the number of sub-directories that will be created when -# CREATE_SUBDIRS tag is set to YES. Level 0 represents 16 directories, and every -# level increment doubles the number of directories, resulting in 4096 -# directories at level 8 which is the default and also the maximum value. The -# sub-directories are organized in 2 levels, the first level always has a fixed -# number of 16 directories. -# Minimum value: 0, maximum value: 8, default value: 8. -# This tag requires that the tag CREATE_SUBDIRS is set to YES. - -CREATE_SUBDIRS_LEVEL = 8 - # If the ALLOW_UNICODE_NAMES tag is set to YES, doxygen will allow non-ASCII # characters to appear in the names of generated files. If set to NO, non-ASCII # characters will be escaped, for example _xE3_x81_x84 will be used for Unicode @@ -109,18 +81,26 @@ ALLOW_UNICODE_NAMES = NO # The OUTPUT_LANGUAGE tag is used to specify the language in which all # documentation generated by doxygen is written. Doxygen will use this # information to generate all constant output in the proper language. -# Possible values are: Afrikaans, Arabic, Armenian, Brazilian, Bulgarian, -# Catalan, Chinese, Chinese-Traditional, Croatian, Czech, Danish, Dutch, English -# (United States), Esperanto, Farsi (Persian), Finnish, French, German, Greek, -# Hindi, Hungarian, Indonesian, Italian, Japanese, Japanese-en (Japanese with -# English messages), Korean, Korean-en (Korean with English messages), Latvian, -# Lithuanian, Macedonian, Norwegian, Persian (Farsi), Polish, Portuguese, -# Romanian, Russian, Serbian, Serbian-Cyrillic, Slovak, Slovene, Spanish, -# Swedish, Turkish, Ukrainian and Vietnamese. +# Possible values are: Afrikaans, Arabic, Armenian, Brazilian, Catalan, Chinese, +# Chinese-Traditional, Croatian, Czech, Danish, Dutch, English (United States), +# Esperanto, Farsi (Persian), Finnish, French, German, Greek, Hungarian, +# Indonesian, Italian, Japanese, Japanese-en (Japanese with English messages), +# Korean, Korean-en (Korean with English messages), Latvian, Lithuanian, +# Macedonian, Norwegian, Persian (Farsi), Polish, Portuguese, Romanian, Russian, +# Serbian, Serbian-Cyrillic, Slovak, Slovene, Spanish, Swedish, Turkish, +# Ukrainian and Vietnamese. # The default value is: English. OUTPUT_LANGUAGE = English +# The OUTPUT_TEXT_DIRECTION tag is used to specify the direction in which all +# documentation generated by doxygen is written. Doxygen will use this +# information to generate all generated output in the proper direction. +# Possible values are: None, LTR, RTL and Context. +# The default value is: None. + +OUTPUT_TEXT_DIRECTION = None + # If the BRIEF_MEMBER_DESC tag is set to YES, doxygen will include brief member # descriptions after the members that are listed in the file and class # documentation (similar to Javadoc). Set to NO to disable this. @@ -278,16 +258,16 @@ TAB_SIZE = 4 # the documentation. An alias has the form: # name=value # For example adding -# "sideeffect=@par Side Effects:^^" +# "sideeffect=@par Side Effects:\n" # will allow you to put the command \sideeffect (or @sideeffect) in the # documentation, which will result in a user-defined paragraph with heading -# "Side Effects:". Note that you cannot put \n's in the value part of an alias -# to insert newlines (in the resulting output). You can put ^^ in the value part -# of an alias to insert a newline as if a physical newline was in the original -# file. When you need a literal { or } or , in the value part of an alias you -# have to escape them by means of a backslash (\), this can lead to conflicts -# with the commands \{ and \} for these it is advised to use the version @{ and -# @} or use a double escape (\\{ and \\}) +# "Side Effects:". You can put \n's in the value part of an alias to insert +# newlines (in the resulting output). You can put ^^ in the value part of an +# alias to insert a newline as if a physical newline was in the original file. +# When you need a literal { or } or , in the value part of an alias you have to +# escape them by means of a backslash (\), this can lead to conflicts with the +# commands \{ and \} for these it is advised to use the version @{ and @} or use +# a double escape (\\{ and \\}) ALIASES = @@ -332,8 +312,8 @@ OPTIMIZE_OUTPUT_SLICE = NO # extension. Doxygen has a built-in mapping, but you can override or extend it # using this tag. The format is ext=language, where ext is a file extension, and # language is one of the parsers supported by doxygen: IDL, Java, JavaScript, -# Csharp (C#), C, C++, Lex, D, PHP, md (Markdown), Objective-C, Python, Slice, -# VHDL, Fortran (fixed format Fortran: FortranFixed, free formatted Fortran: +# Csharp (C#), C, C++, D, PHP, md (Markdown), Objective-C, Python, Slice, VHDL, +# Fortran (fixed format Fortran: FortranFixed, free formatted Fortran: # FortranFree, unknown formatted Fortran: Fortran. In the later case the parser # tries to guess whether the code is fixed or free formatted code, this is the # default for Fortran type files). For instance to make doxygen treat .inc files @@ -369,17 +349,6 @@ MARKDOWN_SUPPORT = YES TOC_INCLUDE_HEADINGS = 5 -# The MARKDOWN_ID_STYLE tag can be used to specify the algorithm used to -# generate identifiers for the Markdown headings. Note: Every identifier is -# unique. -# Possible values are: DOXYGEN use a fixed 'autotoc_md' string followed by a -# sequence number starting at 0 and GITHUB use the lower case version of title -# with any whitespace replaced by '-' and punctuation characters removed. -# The default value is: DOXYGEN. -# This tag requires that the tag MARKDOWN_SUPPORT is set to YES. - -MARKDOWN_ID_STYLE = DOXYGEN - # When enabled doxygen tries to link words that correspond to documented # classes, or namespaces to their corresponding documentation. Such a link can # be prevented in individual cases by putting a % sign in front of the word or @@ -491,27 +460,19 @@ TYPEDEF_HIDES_STRUCT = NO LOOKUP_CACHE_SIZE = 0 -# The NUM_PROC_THREADS specifies the number of threads doxygen is allowed to use +# The NUM_PROC_THREADS specifies the number threads doxygen is allowed to use # during processing. When set to 0 doxygen will based this on the number of # cores available in the system. You can set it explicitly to a value larger # than 0 to get more control over the balance between CPU load and processing # speed. At this moment only the input processing can be done using multiple # threads. Since this is still an experimental feature the default is set to 1, -# which effectively disables parallel processing. Please report any issues you +# which efficively disables parallel processing. Please report any issues you # encounter. Generating dot graphs in parallel is controlled by the # DOT_NUM_THREADS setting. # Minimum value: 0, maximum value: 32, default value: 1. NUM_PROC_THREADS = 1 -# If the TIMESTAMP tag is set different from NO then each generated page will -# contain the date or date and time when the page was generated. Setting this to -# NO can help when comparing the output of multiple runs. -# Possible values are: YES, NO, DATETIME and DATE. -# The default value is: NO. - -TIMESTAMP = NO - #--------------------------------------------------------------------------- # Build related configuration options #--------------------------------------------------------------------------- @@ -588,16 +549,15 @@ RESOLVE_UNNAMED_PARAMS = YES # section is generated. This option has no effect if EXTRACT_ALL is enabled. # The default value is: NO. -HIDE_UNDOC_MEMBERS = NO +HIDE_UNDOC_MEMBERS = YES # If the HIDE_UNDOC_CLASSES tag is set to YES, doxygen will hide all # undocumented classes that are normally visible in the class hierarchy. If set # to NO, these classes will be included in the various overviews. This option -# will also hide undocumented C++ concepts if enabled. This option has no effect -# if EXTRACT_ALL is enabled. +# has no effect if EXTRACT_ALL is enabled. # The default value is: NO. -HIDE_UNDOC_CLASSES = NO +HIDE_UNDOC_CLASSES = YES # If the HIDE_FRIEND_COMPOUNDS tag is set to YES, doxygen will hide all friend # declarations. If set to NO, these declarations will be included in the @@ -625,17 +585,16 @@ INTERNAL_DOCS = NO # filesystem is case sensitive (i.e. it supports files in the same directory # whose names only differ in casing), the option must be set to YES to properly # deal with such files in case they appear in the input. For filesystems that -# are not case sensitive the option should be set to NO to properly deal with +# are not case sensitive the option should be be set to NO to properly deal with # output files written for symbols that only differ in casing, such as for two # classes, one named CLASS and the other named Class, and to also support # references to files without having to specify the exact matching casing. On # Windows (including Cygwin) and MacOS, users should typically set this option # to NO, whereas on Linux or other Unix flavors it should typically be set to # YES. -# Possible values are: SYSTEM, NO and YES. -# The default value is: SYSTEM. +# The default value is: system dependent. -CASE_SENSE_NAMES = SYSTEM +CASE_SENSE_NAMES = YES # If the HIDE_SCOPE_NAMES tag is set to NO then doxygen will show members with # their full class and namespace scopes in the documentation. If set to YES, the @@ -651,12 +610,6 @@ HIDE_SCOPE_NAMES = NO HIDE_COMPOUND_REFERENCE= NO -# If the SHOW_HEADERFILE tag is set to YES then the documentation for a class -# will show which file needs to be included to use the class. -# The default value is: YES. - -SHOW_HEADERFILE = YES - # If the SHOW_INCLUDE_FILES tag is set to YES then doxygen will put a list of # the files that are included by a file in the documentation of that file. # The default value is: YES. @@ -814,8 +767,7 @@ FILE_VERSION_FILTER = # output files in an output format independent way. To create the layout file # that represents doxygen's defaults, run doxygen with the -l option. You can # optionally specify a file name after the option, if omitted DoxygenLayout.xml -# will be used as the name of the layout file. See also section "Changing the -# layout of pages" for information. +# will be used as the name of the layout file. # # Note that if you run doxygen from a directory containing a file called # DoxygenLayout.xml, doxygen will parse it automatically even if the LAYOUT_FILE @@ -861,50 +813,27 @@ WARNINGS = YES WARN_IF_UNDOCUMENTED = YES # If the WARN_IF_DOC_ERROR tag is set to YES, doxygen will generate warnings for -# potential errors in the documentation, such as documenting some parameters in -# a documented function twice, or documenting parameters that don't exist or -# using markup commands wrongly. +# potential errors in the documentation, such as not documenting some parameters +# in a documented function, or documenting parameters that don't exist or using +# markup commands wrongly. # The default value is: YES. WARN_IF_DOC_ERROR = YES -# If WARN_IF_INCOMPLETE_DOC is set to YES, doxygen will warn about incomplete -# function parameter documentation. If set to NO, doxygen will accept that some -# parameters have no documentation without warning. -# The default value is: YES. - -WARN_IF_INCOMPLETE_DOC = YES - # This WARN_NO_PARAMDOC option can be enabled to get warnings for functions that # are documented, but have no documentation for their parameters or return -# value. If set to NO, doxygen will only warn about wrong parameter -# documentation, but not about the absence of documentation. If EXTRACT_ALL is -# set to YES then this flag will automatically be disabled. See also -# WARN_IF_INCOMPLETE_DOC +# value. If set to NO, doxygen will only warn about wrong or incomplete +# parameter documentation, but not about the absence of documentation. If +# EXTRACT_ALL is set to YES then this flag will automatically be disabled. # The default value is: NO. WARN_NO_PARAMDOC = NO -# If WARN_IF_UNDOC_ENUM_VAL option is set to YES, doxygen will warn about -# undocumented enumeration values. If set to NO, doxygen will accept -# undocumented enumeration values. If EXTRACT_ALL is set to YES then this flag -# will automatically be disabled. -# The default value is: NO. - -WARN_IF_UNDOC_ENUM_VAL = NO - # If the WARN_AS_ERROR tag is set to YES then doxygen will immediately stop when # a warning is encountered. If the WARN_AS_ERROR tag is set to FAIL_ON_WARNINGS # then doxygen will continue running as if WARN_AS_ERROR tag is set to NO, but # at the end of the doxygen process doxygen will return with a non-zero status. -# If the WARN_AS_ERROR tag is set to FAIL_ON_WARNINGS_PRINT then doxygen behaves -# like FAIL_ON_WARNINGS but in case no WARN_LOGFILE is defined doxygen will not -# write the warning messages in between other messages but write them at the end -# of a run, in case a WARN_LOGFILE is defined the warning messages will be -# besides being in the defined file also be shown at the end of a run, unless -# the WARN_LOGFILE is defined as - i.e. standard output (stdout) in that case -# the behavior will remain as with the setting FAIL_ON_WARNINGS. -# Possible values are: NO, YES, FAIL_ON_WARNINGS and FAIL_ON_WARNINGS_PRINT. +# Possible values are: NO, YES and FAIL_ON_WARNINGS. # The default value is: NO. WARN_AS_ERROR = NO @@ -915,27 +844,13 @@ WARN_AS_ERROR = NO # and the warning text. Optionally the format may contain $version, which will # be replaced by the version of the file (if it could be obtained via # FILE_VERSION_FILTER) -# See also: WARN_LINE_FORMAT # The default value is: $file:$line: $text. WARN_FORMAT = "$file:$line: $text" -# In the $text part of the WARN_FORMAT command it is possible that a reference -# to a more specific place is given. To make it easier to jump to this place -# (outside of doxygen) the user can define a custom "cut" / "paste" string. -# Example: -# WARN_LINE_FORMAT = "'vi $file +$line'" -# See also: WARN_FORMAT -# The default value is: at line $line of file $file. - -WARN_LINE_FORMAT = "at line $line of file $file" - # The WARN_LOGFILE tag can be used to specify a file to which warning and error # messages should be written. If left blank the output is written to standard -# error (stderr). In case the file specified cannot be opened for writing the -# warning and error messages are written to standard error. When as file - is -# specified the warning and error messages are written to standard output -# (stdout). +# error (stderr). WARN_LOGFILE = @@ -949,28 +864,18 @@ WARN_LOGFILE = # spaces. See also FILE_PATTERNS and EXTENSION_MAPPING # Note: If this tag is empty the current directory is searched. -INPUT = ../include +INPUT = ../include \ + pages # This tag can be used to specify the character encoding of the source files # that doxygen parses. Internally doxygen uses the UTF-8 encoding. Doxygen uses # libiconv (or the iconv built into libc) for the transcoding. See the libiconv # documentation (see: # https://www.gnu.org/software/libiconv/) for the list of possible encodings. -# See also: INPUT_FILE_ENCODING # The default value is: UTF-8. INPUT_ENCODING = UTF-8 -# This tag can be used to specify the character encoding of the source files -# that doxygen parses The INPUT_FILE_ENCODING tag can be used to specify -# character encoding on a per file pattern basis. Doxygen will compare the file -# name with each pattern and apply the encoding instead of the default -# INPUT_ENCODING) if there is a match. The character encodings are a list of the -# form: pattern=encoding (like *.php=ISO-8859-1). See cfg_input_encoding -# "INPUT_ENCODING" for further information on supported encodings. - -INPUT_FILE_ENCODING = - # If the value of the INPUT tag contains directories, you can use the # FILE_PATTERNS tag to specify one or more wildcard patterns (like *.cpp and # *.h) to filter out the source-files in the directories. @@ -982,22 +887,18 @@ INPUT_FILE_ENCODING = # Note the list of default checked file patterns might differ from the list of # default file extension mappings. # -# If left blank the following patterns are tested:*.c, *.cc, *.cxx, *.cxxm, -# *.cpp, *.cppm, *.ccm, *.c++, *.c++m, *.java, *.ii, *.ixx, *.ipp, *.i++, *.inl, -# *.idl, *.ddl, *.odl, *.h, *.hh, *.hxx, *.hpp, *.h++, *.ixx, *.l, *.cs, *.d, -# *.php, *.php4, *.php5, *.phtml, *.inc, *.m, *.markdown, *.md, *.mm, *.dox (to -# be provided as doxygen C comment), *.py, *.pyw, *.f90, *.f95, *.f03, *.f08, -# *.f18, *.f, *.for, *.vhd, *.vhdl, *.ucf, *.qsf and *.ice. +# If left blank the following patterns are tested:*.c, *.cc, *.cxx, *.cpp, +# *.c++, *.java, *.ii, *.ixx, *.ipp, *.i++, *.inl, *.idl, *.ddl, *.odl, *.h, +# *.hh, *.hxx, *.hpp, *.h++, *.cs, *.d, *.php, *.php4, *.php5, *.phtml, *.inc, +# *.m, *.markdown, *.md, *.mm, *.dox (to be provided as doxygen C comment), +# *.py, *.pyw, *.f90, *.f95, *.f03, *.f08, *.f18, *.f, *.for, *.vhd, *.vhdl, +# *.ucf, *.qsf and *.ice. FILE_PATTERNS = *.c \ *.cc \ *.cxx \ - *.cxxm \ *.cpp \ - *.cppm \ - *.ccm \ *.c++ \ - *.c++m \ *.java \ *.ii \ *.ixx \ @@ -1012,8 +913,6 @@ FILE_PATTERNS = *.c \ *.hxx \ *.hpp \ *.h++ \ - *.ixx \ - *.l \ *.cs \ *.d \ *.php \ @@ -1076,7 +975,10 @@ EXCLUDE_PATTERNS = # (namespaces, classes, functions, etc.) that should be excluded from the # output. The symbol name can be a fully qualified name, a word, or if the # wildcard * is used, a substring. Examples: ANamespace, AClass, -# ANamespace::AClass, ANamespace::*Test +# AClass::ANamespace, ANamespace::*Test +# +# Note that the wildcards are matched against the file with absolute path, so to +# exclude all test directories use the pattern */test/* EXCLUDE_SYMBOLS = @@ -1121,11 +1023,6 @@ IMAGE_PATH = # code is scanned, but not when the output code is generated. If lines are added # or removed, the anchors will not be placed correctly. # -# Note that doxygen will use the data processed and written to standard output -# for further processing, therefore nothing else, like debug statements or used -# commands (so in case of a Windows batch file always use @echo OFF), should be -# written to standard output. -# # Note that for custom extensions or not directly supported extensions you also # need to set EXTENSION_MAPPING for the extension otherwise the files are not # properly processed by doxygen. @@ -1167,15 +1064,6 @@ FILTER_SOURCE_PATTERNS = USE_MDFILE_AS_MAINPAGE = -# The Fortran standard specifies that for fixed formatted Fortran code all -# characters from position 72 are to be considered as comment. A common -# extension is to allow longer lines before the automatic comment starts. The -# setting FORTRAN_COMMENT_AFTER will also make it possible that longer lines can -# be processed before the automatic comment starts. -# Minimum value: 7, maximum value: 10000, default value: 72. - -FORTRAN_COMMENT_AFTER = 72 - #--------------------------------------------------------------------------- # Configuration options related to source browsing #--------------------------------------------------------------------------- @@ -1190,8 +1078,7 @@ FORTRAN_COMMENT_AFTER = 72 SOURCE_BROWSER = NO # Setting the INLINE_SOURCES tag to YES will include the body of functions, -# multi-line macros, enums or list initialized variables directly into the -# documentation. +# classes and enums directly into the documentation. # The default value is: NO. INLINE_SOURCES = NO @@ -1263,6 +1150,44 @@ USE_HTAGS = NO VERBATIM_HEADERS = YES +# If the CLANG_ASSISTED_PARSING tag is set to YES then doxygen will use the +# clang parser (see: +# http://clang.llvm.org/) for more accurate parsing at the cost of reduced +# performance. This can be particularly helpful with template rich C++ code for +# which doxygen's built-in parser lacks the necessary type information. +# Note: The availability of this option depends on whether or not doxygen was +# generated with the -Duse_libclang=ON option for CMake. +# The default value is: NO. + +CLANG_ASSISTED_PARSING = NO + +# If clang assisted parsing is enabled and the CLANG_ADD_INC_PATHS tag is set to +# YES then doxygen will add the directory of each input to the include path. +# The default value is: YES. + +CLANG_ADD_INC_PATHS = YES + +# If clang assisted parsing is enabled you can provide the compiler with command +# line options that you would normally use when invoking the compiler. Note that +# the include paths will already be set by doxygen for the files and directories +# specified with INPUT and INCLUDE_PATH. +# This tag requires that the tag CLANG_ASSISTED_PARSING is set to YES. + +CLANG_OPTIONS = + +# If clang assisted parsing is enabled you can provide the clang parser with the +# path to the directory containing a file called compile_commands.json. This +# file is the compilation database (see: +# http://clang.llvm.org/docs/HowToSetupToolingForLLVM.html) containing the +# options used when the source files were built. This is equivalent to +# specifying the -p option to a clang tool, such as clang-check. These options +# will then be passed to the parser. Any options specified with CLANG_OPTIONS +# will be added as well. +# Note: The availability of this option depends on whether or not doxygen was +# generated with the -Duse_libclang=ON option for CMake. + +CLANG_DATABASE_PATH = + #--------------------------------------------------------------------------- # Configuration options related to the alphabetical class index #--------------------------------------------------------------------------- @@ -1274,11 +1199,10 @@ VERBATIM_HEADERS = YES ALPHABETICAL_INDEX = YES -# The IGNORE_PREFIX tag can be used to specify a prefix (or a list of prefixes) -# that should be ignored while generating the index headers. The IGNORE_PREFIX -# tag works for classes, function and member names. The entity will be placed in -# the alphabetical list under the first letter of the entity name that remains -# after removing the prefix. +# In case all classes in a project start with a common prefix, all classes will +# be put under the same header in the alphabetical index. The IGNORE_PREFIX tag +# can be used to specify a prefix (or a list of prefixes) that should be ignored +# while generating the index headers. # This tag requires that the tag ALPHABETICAL_INDEX is set to YES. IGNORE_PREFIX = @@ -1290,7 +1214,7 @@ IGNORE_PREFIX = # If the GENERATE_HTML tag is set to YES, doxygen will generate HTML output # The default value is: YES. -GENERATE_HTML = NO +GENERATE_HTML = YES # The HTML_OUTPUT tag is used to specify where the HTML docs will be put. If a # relative path is entered the value of OUTPUT_DIRECTORY will be put in front of @@ -1357,15 +1281,11 @@ HTML_STYLESHEET = # Doxygen will copy the style sheet files to the output directory. # Note: The order of the extra style sheet files is of importance (e.g. the last # style sheet in the list overrules the setting of the previous ones in the -# list). -# Note: Since the styling of scrollbars can currently not be overruled in -# Webkit/Chromium, the styling will be left out of the default doxygen.css if -# one or more extra stylesheets have been specified. So if scrollbar -# customization is desired it has to be added explicitly. For an example see the -# documentation. +# list). For an example see the documentation. # This tag requires that the tag GENERATE_HTML is set to YES. -HTML_EXTRA_STYLESHEET = +HTML_EXTRA_STYLESHEET = doxygen-awesome-css/doxygen-awesome.css \ + doxygen-awesome-css/doxygen-awesome-sidebar-only.css # The HTML_EXTRA_FILES tag can be used to specify one or more extra images or # other source files which should be copied to the HTML output directory. Note @@ -1377,22 +1297,9 @@ HTML_EXTRA_STYLESHEET = HTML_EXTRA_FILES = -# The HTML_COLORSTYLE tag can be used to specify if the generated HTML output -# should be rendered with a dark or light theme. -# Possible values are: LIGHT always generate light mode output, DARK always -# generate dark mode output, AUTO_LIGHT automatically set the mode according to -# the user preference, use light mode if no preference is set (the default), -# AUTO_DARK automatically set the mode according to the user preference, use -# dark mode if no preference is set and TOGGLE allow to user to switch between -# light and dark mode via a button. -# The default value is: AUTO_LIGHT. -# This tag requires that the tag GENERATE_HTML is set to YES. - -HTML_COLORSTYLE = AUTO_LIGHT - # The HTML_COLORSTYLE_HUE tag controls the color of the HTML output. Doxygen # will adjust the colors in the style sheet and background images according to -# this color. Hue is specified as an angle on a color-wheel, see +# this color. Hue is specified as an angle on a colorwheel, see # https://en.wikipedia.org/wiki/Hue for more information. For instance the value # 0 represents red, 60 is yellow, 120 is green, 180 is cyan, 240 is blue, 300 # purple, and 360 is red again. @@ -1402,7 +1309,7 @@ HTML_COLORSTYLE = AUTO_LIGHT HTML_COLORSTYLE_HUE = 220 # The HTML_COLORSTYLE_SAT tag controls the purity (or saturation) of the colors -# in the HTML output. For a value of 0 the output will use gray-scales only. A +# in the HTML output. For a value of 0 the output will use grayscales only. A # value of 255 will produce the most vivid colors. # Minimum value: 0, maximum value: 255, default value: 100. # This tag requires that the tag GENERATE_HTML is set to YES. @@ -1420,6 +1327,15 @@ HTML_COLORSTYLE_SAT = 100 HTML_COLORSTYLE_GAMMA = 80 +# If the HTML_TIMESTAMP tag is set to YES then the footer of each generated HTML +# page will contain the date and time when the page was generated. Setting this +# to YES can help to show when doxygen was last run and thus if the +# documentation is up to date. +# The default value is: NO. +# This tag requires that the tag GENERATE_HTML is set to YES. + +HTML_TIMESTAMP = NO + # If the HTML_DYNAMIC_MENUS tag is set to YES then the generated HTML # documentation will contain a main index with vertical navigation menus that # are dynamically created via JavaScript. If disabled, the navigation index will @@ -1439,33 +1355,6 @@ HTML_DYNAMIC_MENUS = YES HTML_DYNAMIC_SECTIONS = NO -# If the HTML_CODE_FOLDING tag is set to YES then classes and functions can be -# dynamically folded and expanded in the generated HTML source code. -# The default value is: YES. -# This tag requires that the tag GENERATE_HTML is set to YES. - -HTML_CODE_FOLDING = YES - -# If the HTML_COPY_CLIPBOARD tag is set to YES then doxygen will show an icon in -# the top right corner of code and text fragments that allows the user to copy -# its content to the clipboard. Note this only works if supported by the browser -# and the web page is served via a secure context (see: -# https://www.w3.org/TR/secure-contexts/), i.e. using the https: or file: -# protocol. -# The default value is: YES. -# This tag requires that the tag GENERATE_HTML is set to YES. - -HTML_COPY_CLIPBOARD = YES - -# Doxygen stores a couple of settings persistently in the browser (via e.g. -# cookies). By default these settings apply to all HTML pages generated by -# doxygen across all projects. The HTML_PROJECT_COOKIE tag can be used to store -# the settings under a project specific key, such that the user preferences will -# be stored separately. -# This tag requires that the tag GENERATE_HTML is set to YES. - -HTML_PROJECT_COOKIE = - # With HTML_INDEX_NUM_ENTRIES one can control the preferred number of entries # shown in the various tree structured indices initially; the user can expand # and collapse entries dynamically later on. Doxygen will expand the tree to @@ -1502,13 +1391,6 @@ GENERATE_DOCSET = NO DOCSET_FEEDNAME = "Doxygen generated docs" -# This tag determines the URL of the docset feed. A documentation feed provides -# an umbrella under which multiple documentation sets from a single provider -# (such as a company or product suite) can be grouped. -# This tag requires that the tag GENERATE_DOCSET is set to YES. - -DOCSET_FEEDURL = - # This tag specifies a string that should uniquely identify the documentation # set bundle. This should be a reverse domain-name style string, e.g. # com.mycompany.MyDocSet. Doxygen will append .docset to the name. @@ -1534,12 +1416,8 @@ DOCSET_PUBLISHER_NAME = Publisher # If the GENERATE_HTMLHELP tag is set to YES then doxygen generates three # additional HTML index files: index.hhp, index.hhc, and index.hhk. The # index.hhp is a project file that can be read by Microsoft's HTML Help Workshop -# on Windows. In the beginning of 2021 Microsoft took the original page, with -# a.o. the download links, offline the HTML help workshop was already many years -# in maintenance mode). You can download the HTML help workshop from the web -# archives at Installation executable (see: -# http://web.archive.org/web/20160201063255/http://download.microsoft.com/downlo -# ad/0/A/9/0A939EF6-E31C-430F-A3DF-DFAE7960D564/htmlhelp.exe). +# (see: +# https://www.microsoft.com/en-us/download/details.aspx?id=21138) on Windows. # # The HTML Help Workshop contains a compiler that can convert all HTML output # generated by doxygen into a single compiled HTML file (.chm). Compiled HTML @@ -1596,16 +1474,6 @@ BINARY_TOC = NO TOC_EXPAND = NO -# The SITEMAP_URL tag is used to specify the full URL of the place where the -# generated documentation will be placed on the server by the user during the -# deployment of the documentation. The generated sitemap is called sitemap.xml -# and placed on the directory specified by HTML_OUTPUT. In case no SITEMAP_URL -# is specified no sitemap is generated. For information about the sitemap -# protocol see https://www.sitemaps.org -# This tag requires that the tag GENERATE_HTML is set to YES. - -SITEMAP_URL = - # If the GENERATE_QHP tag is set to YES and both QHP_NAMESPACE and # QHP_VIRTUAL_FOLDER are set, an additional index file will be generated that # can be used as input for Qt's qhelpgenerator to generate a Qt Compressed Help @@ -1699,7 +1567,7 @@ ECLIPSE_DOC_ID = org.doxygen.Project # The default value is: NO. # This tag requires that the tag GENERATE_HTML is set to YES. -DISABLE_INDEX = NO +DISABLE_INDEX = YES # The GENERATE_TREEVIEW tag is used to specify whether a tree-like index # structure should be generated to display hierarchical information. If the tag @@ -1708,27 +1576,15 @@ DISABLE_INDEX = NO # to work a browser that supports JavaScript, DHTML, CSS and frames is required # (i.e. any modern browser). Windows users are probably better off using the # HTML help feature. Via custom style sheets (see HTML_EXTRA_STYLESHEET) one can -# further fine tune the look of the index (see "Fine-tuning the output"). As an -# example, the default style sheet generated by doxygen has an example that -# shows how to put an image at the root of the tree instead of the PROJECT_NAME. -# Since the tree basically has the same information as the tab index, you could -# consider setting DISABLE_INDEX to YES when enabling this option. -# The default value is: NO. -# This tag requires that the tag GENERATE_HTML is set to YES. - -GENERATE_TREEVIEW = NO - -# When both GENERATE_TREEVIEW and DISABLE_INDEX are set to YES, then the -# FULL_SIDEBAR option determines if the side bar is limited to only the treeview -# area (value NO) or if it should extend to the full height of the window (value -# YES). Setting this to YES gives a layout similar to -# https://docs.readthedocs.io with more room for contents, but less room for the -# project logo, title, and description. If either GENERATE_TREEVIEW or -# DISABLE_INDEX is set to NO, this option has no effect. +# further fine-tune the look of the index. As an example, the default style +# sheet generated by doxygen has an example that shows how to put an image at +# the root of the tree instead of the PROJECT_NAME. Since the tree basically has +# the same information as the tab index, you could consider setting +# DISABLE_INDEX to YES when enabling this option. # The default value is: NO. # This tag requires that the tag GENERATE_HTML is set to YES. -FULL_SIDEBAR = NO +GENERATE_TREEVIEW = YES # The ENUM_VALUES_PER_LINE tag can be used to set the number of enum values that # doxygen will group on one line in the generated HTML documentation. @@ -1754,13 +1610,6 @@ TREEVIEW_WIDTH = 250 EXT_LINKS_IN_WINDOW = NO -# If the OBFUSCATE_EMAILS tag is set to YES, doxygen will obfuscate email -# addresses. -# The default value is: YES. -# This tag requires that the tag GENERATE_HTML is set to YES. - -OBFUSCATE_EMAILS = YES - # If the HTML_FORMULA_FORMAT option is set to svg, doxygen will use the pdf2svg # tool (see https://github.com/dawbarton/pdf2svg) or inkscape (see # https://inkscape.org) to generate formulas as SVG images instead of PNGs for @@ -1781,6 +1630,17 @@ HTML_FORMULA_FORMAT = png FORMULA_FONTSIZE = 10 +# Use the FORMULA_TRANSPARENT tag to determine whether or not the images +# generated for formulas are transparent PNGs. Transparent PNGs are not +# supported properly for IE 6.0, but are supported on all modern browsers. +# +# Note that when changing this option you need to delete any form_*.png files in +# the HTML output directory before the changes have effect. +# The default value is: YES. +# This tag requires that the tag GENERATE_HTML is set to YES. + +FORMULA_TRANSPARENT = YES + # The FORMULA_MACROFILE can contain LaTeX \newcommand and \renewcommand commands # to create new LaTeX commands to be used in formulas as building blocks. See # the section "Including formulas" for details. @@ -1798,29 +1658,11 @@ FORMULA_MACROFILE = USE_MATHJAX = NO -# With MATHJAX_VERSION it is possible to specify the MathJax version to be used. -# Note that the different versions of MathJax have different requirements with -# regards to the different settings, so it is possible that also other MathJax -# settings have to be changed when switching between the different MathJax -# versions. -# Possible values are: MathJax_2 and MathJax_3. -# The default value is: MathJax_2. -# This tag requires that the tag USE_MATHJAX is set to YES. - -MATHJAX_VERSION = MathJax_2 - # When MathJax is enabled you can set the default output format to be used for -# the MathJax output. For more details about the output format see MathJax -# version 2 (see: -# http://docs.mathjax.org/en/v2.7-latest/output.html) and MathJax version 3 -# (see: -# http://docs.mathjax.org/en/latest/web/components/output.html). +# the MathJax output. See the MathJax site (see: +# http://docs.mathjax.org/en/v2.7-latest/output.html) for more details. # Possible values are: HTML-CSS (which is slower, but has the best -# compatibility. This is the name for Mathjax version 2, for MathJax version 3 -# this will be translated into chtml), NativeMML (i.e. MathML. Only supported -# for NathJax 2. For MathJax version 3 chtml will be used instead.), chtml (This -# is the name for Mathjax version 3, for MathJax version 2 this will be -# translated into HTML-CSS) and SVG. +# compatibility), NativeMML (i.e. MathML) and SVG. # The default value is: HTML-CSS. # This tag requires that the tag USE_MATHJAX is set to YES. @@ -1833,21 +1675,15 @@ MATHJAX_FORMAT = HTML-CSS # MATHJAX_RELPATH should be ../mathjax. The default value points to the MathJax # Content Delivery Network so you can quickly see the result without installing # MathJax. However, it is strongly recommended to install a local copy of -# MathJax from https://www.mathjax.org before deployment. The default value is: -# - in case of MathJax version 2: https://cdn.jsdelivr.net/npm/mathjax@2 -# - in case of MathJax version 3: https://cdn.jsdelivr.net/npm/mathjax@3 +# MathJax from https://www.mathjax.org before deployment. +# The default value is: https://cdn.jsdelivr.net/npm/mathjax@2. # This tag requires that the tag USE_MATHJAX is set to YES. -MATHJAX_RELPATH = +MATHJAX_RELPATH = https://cdn.jsdelivr.net/npm/mathjax@2 # The MATHJAX_EXTENSIONS tag can be used to specify one or more MathJax # extension names that should be enabled during MathJax rendering. For example -# for MathJax version 2 (see -# https://docs.mathjax.org/en/v2.7-latest/tex.html#tex-and-latex-extensions): # MATHJAX_EXTENSIONS = TeX/AMSmath TeX/AMSsymbols -# For example for MathJax version 3 (see -# http://docs.mathjax.org/en/latest/input/tex/extensions/index.html): -# MATHJAX_EXTENSIONS = ams # This tag requires that the tag USE_MATHJAX is set to YES. MATHJAX_EXTENSIONS = @@ -2027,31 +1863,29 @@ PAPER_TYPE = a4 EXTRA_PACKAGES = -# The LATEX_HEADER tag can be used to specify a user-defined LaTeX header for -# the generated LaTeX document. The header should contain everything until the -# first chapter. If it is left blank doxygen will generate a standard header. It -# is highly recommended to start with a default header using -# doxygen -w latex new_header.tex new_footer.tex new_stylesheet.sty -# and then modify the file new_header.tex. See also section "Doxygen usage" for -# information on how to generate the default header that doxygen normally uses. +# The LATEX_HEADER tag can be used to specify a personal LaTeX header for the +# generated LaTeX document. The header should contain everything until the first +# chapter. If it is left blank doxygen will generate a standard header. See +# section "Doxygen usage" for information on how to let doxygen write the +# default header to a separate file. # -# Note: Only use a user-defined header if you know what you are doing! -# Note: The header is subject to change so you typically have to regenerate the -# default header when upgrading to a newer version of doxygen. The following -# commands have a special meaning inside the header (and footer): For a -# description of the possible markers and block names see the documentation. +# Note: Only use a user-defined header if you know what you are doing! The +# following commands have a special meaning inside the header: $title, +# $datetime, $date, $doxygenversion, $projectname, $projectnumber, +# $projectbrief, $projectlogo. Doxygen will replace $title with the empty +# string, for the replacement values of the other commands the user is referred +# to HTML_HEADER. # This tag requires that the tag GENERATE_LATEX is set to YES. LATEX_HEADER = -# The LATEX_FOOTER tag can be used to specify a user-defined LaTeX footer for -# the generated LaTeX document. The footer should contain everything after the -# last chapter. If it is left blank doxygen will generate a standard footer. See +# The LATEX_FOOTER tag can be used to specify a personal LaTeX footer for the +# generated LaTeX document. The footer should contain everything after the last +# chapter. If it is left blank doxygen will generate a standard footer. See # LATEX_HEADER for more information on how to generate a default footer and what -# special commands can be used inside the footer. See also section "Doxygen -# usage" for information on how to generate the default footer that doxygen -# normally uses. Note: Only use a user-defined footer if you know what you are -# doing! +# special commands can be used inside the footer. +# +# Note: Only use a user-defined footer if you know what you are doing! # This tag requires that the tag GENERATE_LATEX is set to YES. LATEX_FOOTER = @@ -2094,16 +1928,10 @@ PDF_HYPERLINKS = YES USE_PDFLATEX = YES -# The LATEX_BATCHMODE tag signals the behavior of LaTeX in case of an error. -# Possible values are: NO same as ERROR_STOP, YES same as BATCH, BATCH In batch -# mode nothing is printed on the terminal, errors are scrolled as if is -# hit at every error; missing files that TeX tries to input or request from -# keyboard input (\read on a not open input stream) cause the job to abort, -# NON_STOP In nonstop mode the diagnostic message will appear on the terminal, -# but there is no possibility of user interaction just like in batch mode, -# SCROLL In scroll mode, TeX will stop only for missing files to input or if -# keyboard input is necessary and ERROR_STOP In errorstop mode, TeX will stop at -# each error, asking for user intervention. +# If the LATEX_BATCHMODE tag is set to YES, doxygen will add the \batchmode +# command to the generated LaTeX files. This will instruct LaTeX to keep running +# if errors occur, instead of asking the user for help. This option is also used +# when generating formulas in HTML. # The default value is: NO. # This tag requires that the tag GENERATE_LATEX is set to YES. @@ -2116,6 +1944,16 @@ LATEX_BATCHMODE = NO LATEX_HIDE_INDICES = NO +# If the LATEX_SOURCE_CODE tag is set to YES then doxygen will include source +# code with syntax highlighting in the LaTeX output. +# +# Note that which sources are shown also depends on other settings such as +# SOURCE_BROWSER. +# The default value is: NO. +# This tag requires that the tag GENERATE_LATEX is set to YES. + +LATEX_SOURCE_CODE = NO + # The LATEX_BIB_STYLE tag can be used to specify the style to use for the # bibliography, e.g. plainnat, or ieeetr. See # https://en.wikipedia.org/wiki/BibTeX and \cite for more info. @@ -2124,6 +1962,14 @@ LATEX_HIDE_INDICES = NO LATEX_BIB_STYLE = plain +# If the LATEX_TIMESTAMP tag is set to YES then the footer of each generated +# page will contain the date and time when the page was generated. Setting this +# to NO can help when comparing the output of multiple runs. +# The default value is: NO. +# This tag requires that the tag GENERATE_LATEX is set to YES. + +LATEX_TIMESTAMP = NO + # The LATEX_EMOJI_DIRECTORY tag is used to specify the (relative or absolute) # path from which the emoji images will be read. If a relative path is entered, # it will be relative to the LATEX_OUTPUT directory. If left blank the @@ -2188,6 +2034,16 @@ RTF_STYLESHEET_FILE = RTF_EXTENSIONS_FILE = +# If the RTF_SOURCE_CODE tag is set to YES then doxygen will include source code +# with syntax highlighting in the RTF output. +# +# Note that which sources are shown also depends on other settings such as +# SOURCE_BROWSER. +# The default value is: NO. +# This tag requires that the tag GENERATE_RTF is set to YES. + +RTF_SOURCE_CODE = NO + #--------------------------------------------------------------------------- # Configuration options related to the man page output #--------------------------------------------------------------------------- @@ -2240,7 +2096,7 @@ MAN_LINKS = NO # captures the structure of the code including all documentation. # The default value is: NO. -GENERATE_XML = YES +GENERATE_XML = NO # The XML_OUTPUT tag is used to specify where the XML pages will be put. If a # relative path is entered the value of OUTPUT_DIRECTORY will be put in front of @@ -2284,44 +2140,27 @@ GENERATE_DOCBOOK = NO DOCBOOK_OUTPUT = docbook +# If the DOCBOOK_PROGRAMLISTING tag is set to YES, doxygen will include the +# program listings (including syntax highlighting and cross-referencing +# information) to the DOCBOOK output. Note that enabling this will significantly +# increase the size of the DOCBOOK output. +# The default value is: NO. +# This tag requires that the tag GENERATE_DOCBOOK is set to YES. + +DOCBOOK_PROGRAMLISTING = NO + #--------------------------------------------------------------------------- # Configuration options for the AutoGen Definitions output #--------------------------------------------------------------------------- # If the GENERATE_AUTOGEN_DEF tag is set to YES, doxygen will generate an -# AutoGen Definitions (see https://autogen.sourceforge.net/) file that captures +# AutoGen Definitions (see http://autogen.sourceforge.net/) file that captures # the structure of the code including all documentation. Note that this feature # is still experimental and incomplete at the moment. # The default value is: NO. GENERATE_AUTOGEN_DEF = NO -#--------------------------------------------------------------------------- -# Configuration options related to Sqlite3 output -#--------------------------------------------------------------------------- - -# If the GENERATE_SQLITE3 tag is set to YES doxygen will generate a Sqlite3 -# database with symbols found by doxygen stored in tables. -# The default value is: NO. - -GENERATE_SQLITE3 = NO - -# The SQLITE3_OUTPUT tag is used to specify where the Sqlite3 database will be -# put. If a relative path is entered the value of OUTPUT_DIRECTORY will be put -# in front of it. -# The default directory is: sqlite3. -# This tag requires that the tag GENERATE_SQLITE3 is set to YES. - -SQLITE3_OUTPUT = sqlite3 - -# The SQLITE3_RECREATE_DB tag is set to YES, the existing doxygen_sqlite3.db -# database file will be recreated with each doxygen run. If set to NO, doxygen -# will warn if a database file is already found and not modify it. -# The default value is: YES. -# This tag requires that the tag GENERATE_SQLITE3 is set to YES. - -SQLITE3_RECREATE_DB = YES - #--------------------------------------------------------------------------- # Configuration options related to the Perl module output #--------------------------------------------------------------------------- @@ -2396,8 +2235,7 @@ SEARCH_INCLUDES = YES # The INCLUDE_PATH tag can be used to specify one or more directories that # contain include files that are not input files but should be processed by the -# preprocessor. Note that the INCLUDE_PATH is not recursive, so the setting of -# RECURSIVE has no effect here. +# preprocessor. # This tag requires that the tag SEARCH_INCLUDES is set to YES. INCLUDE_PATH = @@ -2464,15 +2302,15 @@ TAGFILES = GENERATE_TAGFILE = -# If the ALLEXTERNALS tag is set to YES, all external classes and namespaces -# will be listed in the class and namespace index. If set to NO, only the -# inherited external classes will be listed. +# If the ALLEXTERNALS tag is set to YES, all external class will be listed in +# the class index. If set to NO, only the inherited external classes will be +# listed. # The default value is: NO. ALLEXTERNALS = NO # If the EXTERNAL_GROUPS tag is set to YES, all external groups will be listed -# in the topic index. If set to NO, only the current project's groups will be +# in the modules index. If set to NO, only the current project's groups will be # listed. # The default value is: YES. @@ -2486,9 +2324,25 @@ EXTERNAL_GROUPS = YES EXTERNAL_PAGES = YES #--------------------------------------------------------------------------- -# Configuration options related to diagram generator tools +# Configuration options related to the dot tool #--------------------------------------------------------------------------- +# If the CLASS_DIAGRAMS tag is set to YES, doxygen will generate a class diagram +# (in HTML and LaTeX) for classes with base or super classes. Setting the tag to +# NO turns the diagrams off. Note that this option also works with HAVE_DOT +# disabled, but it is recommended to install and use dot, since it yields more +# powerful graphs. +# The default value is: YES. + +CLASS_DIAGRAMS = YES + +# You can include diagrams made with dia in doxygen documentation. Doxygen will +# then run dia to produce the diagram and insert it in the documentation. The +# DIA_PATH tag allows you to specify the directory where the dia binary resides. +# If left empty dia is assumed to be found in the default search path. + +DIA_PATH = + # If set to YES the inheritance and collaboration graphs will hide inheritance # and usage relations if the target is undocumented or is not a class. # The default value is: YES. @@ -2497,12 +2351,12 @@ HIDE_UNDOC_RELATIONS = YES # If you set the HAVE_DOT tag to YES then doxygen will assume the dot tool is # available from the path. This tool is part of Graphviz (see: -# https://www.graphviz.org/), a graph visualization toolkit from AT&T and Lucent +# http://www.graphviz.org/), a graph visualization toolkit from AT&T and Lucent # Bell Labs. The other options in this section have no effect if this option is # set to NO -# The default value is: NO. +# The default value is: YES. -HAVE_DOT = NO +HAVE_DOT = YES # The DOT_NUM_THREADS specifies the number of dot invocations doxygen is allowed # to run in parallel. When set to 0 doxygen will base this on the number of @@ -2514,77 +2368,49 @@ HAVE_DOT = NO DOT_NUM_THREADS = 0 -# DOT_COMMON_ATTR is common attributes for nodes, edges and labels of -# subgraphs. When you want a differently looking font in the dot files that -# doxygen generates you can specify fontname, fontcolor and fontsize attributes. -# For details please see Node, -# Edge and Graph Attributes specification You need to make sure dot is able -# to find the font, which can be done by putting it in a standard location or by -# setting the DOTFONTPATH environment variable or by setting DOT_FONTPATH to the -# directory containing the font. Default graphviz fontsize is 14. -# The default value is: fontname=Helvetica,fontsize=10. +# When you want a differently looking font in the dot files that doxygen +# generates you can specify the font name using DOT_FONTNAME. You need to make +# sure dot is able to find the font, which can be done by putting it in a +# standard location or by setting the DOTFONTPATH environment variable or by +# setting DOT_FONTPATH to the directory containing the font. +# The default value is: Helvetica. # This tag requires that the tag HAVE_DOT is set to YES. -DOT_COMMON_ATTR = "fontname=Helvetica,fontsize=10" +DOT_FONTNAME = Helvetica -# DOT_EDGE_ATTR is concatenated with DOT_COMMON_ATTR. For elegant style you can -# add 'arrowhead=open, arrowtail=open, arrowsize=0.5'. Complete documentation about -# arrows shapes. -# The default value is: labelfontname=Helvetica,labelfontsize=10. +# The DOT_FONTSIZE tag can be used to set the size (in points) of the font of +# dot graphs. +# Minimum value: 4, maximum value: 24, default value: 10. # This tag requires that the tag HAVE_DOT is set to YES. -DOT_EDGE_ATTR = "labelfontname=Helvetica,labelfontsize=10" +DOT_FONTSIZE = 10 -# DOT_NODE_ATTR is concatenated with DOT_COMMON_ATTR. For view without boxes -# around nodes set 'shape=plain' or 'shape=plaintext' Shapes specification -# The default value is: shape=box,height=0.2,width=0.4. -# This tag requires that the tag HAVE_DOT is set to YES. - -DOT_NODE_ATTR = "shape=box,height=0.2,width=0.4" - -# You can set the path where dot can find font specified with fontname in -# DOT_COMMON_ATTR and others dot attributes. +# By default doxygen will tell dot to use the default font as specified with +# DOT_FONTNAME. If you specify a different font using DOT_FONTNAME you can set +# the path where dot can find it using this tag. # This tag requires that the tag HAVE_DOT is set to YES. DOT_FONTPATH = -# If the CLASS_GRAPH tag is set to YES or GRAPH or BUILTIN then doxygen will -# generate a graph for each documented class showing the direct and indirect -# inheritance relations. In case the CLASS_GRAPH tag is set to YES or GRAPH and -# HAVE_DOT is enabled as well, then dot will be used to draw the graph. In case -# the CLASS_GRAPH tag is set to YES and HAVE_DOT is disabled or if the -# CLASS_GRAPH tag is set to BUILTIN, then the built-in generator will be used. -# If the CLASS_GRAPH tag is set to TEXT the direct and indirect inheritance -# relations will be shown as texts / links. Explicit enabling an inheritance -# graph or choosing a different representation for an inheritance graph of a -# specific class, can be accomplished by means of the command \inheritancegraph. -# Disabling an inheritance graph can be accomplished by means of the command -# \hideinheritancegraph. -# Possible values are: NO, YES, TEXT, GRAPH and BUILTIN. +# If the CLASS_GRAPH tag is set to YES then doxygen will generate a graph for +# each documented class showing the direct and indirect inheritance relations. +# Setting this tag to YES will force the CLASS_DIAGRAMS tag to NO. # The default value is: YES. +# This tag requires that the tag HAVE_DOT is set to YES. CLASS_GRAPH = YES # If the COLLABORATION_GRAPH tag is set to YES then doxygen will generate a # graph for each documented class showing the direct and indirect implementation # dependencies (inheritance, containment, and class references variables) of the -# class with other documented classes. Explicit enabling a collaboration graph, -# when COLLABORATION_GRAPH is set to NO, can be accomplished by means of the -# command \collaborationgraph. Disabling a collaboration graph can be -# accomplished by means of the command \hidecollaborationgraph. +# class with other documented classes. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. COLLABORATION_GRAPH = YES # If the GROUP_GRAPHS tag is set to YES then doxygen will generate a graph for -# groups, showing the direct groups dependencies. Explicit enabling a group -# dependency graph, when GROUP_GRAPHS is set to NO, can be accomplished by means -# of the command \groupgraph. Disabling a directory graph can be accomplished by -# means of the command \hidegroupgraph. See also the chapter Grouping in the -# manual. +# groups, showing the direct groups dependencies. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. @@ -2626,8 +2452,8 @@ DOT_UML_DETAILS = NO # The DOT_WRAP_THRESHOLD tag can be used to set the maximum number of characters # to display on a single line. If the actual line length exceeds this threshold -# significantly it will be wrapped across multiple lines. Some heuristics are -# applied to avoid ugly line breaks. +# significantly it will wrapped across multiple lines. Some heuristics are apply +# to avoid ugly line breaks. # Minimum value: 0, maximum value: 1000, default value: 17. # This tag requires that the tag HAVE_DOT is set to YES. @@ -2644,9 +2470,7 @@ TEMPLATE_RELATIONS = NO # If the INCLUDE_GRAPH, ENABLE_PREPROCESSING and SEARCH_INCLUDES tags are set to # YES then doxygen will generate a graph for each documented file showing the # direct and indirect include dependencies of the file with other documented -# files. Explicit enabling an include graph, when INCLUDE_GRAPH is is set to NO, -# can be accomplished by means of the command \includegraph. Disabling an -# include graph can be accomplished by means of the command \hideincludegraph. +# files. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. @@ -2655,10 +2479,7 @@ INCLUDE_GRAPH = YES # If the INCLUDED_BY_GRAPH, ENABLE_PREPROCESSING and SEARCH_INCLUDES tags are # set to YES then doxygen will generate a graph for each documented file showing # the direct and indirect include dependencies of the file with other documented -# files. Explicit enabling an included by graph, when INCLUDED_BY_GRAPH is set -# to NO, can be accomplished by means of the command \includedbygraph. Disabling -# an included by graph can be accomplished by means of the command -# \hideincludedbygraph. +# files. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. @@ -2698,30 +2519,22 @@ GRAPHICAL_HIERARCHY = YES # If the DIRECTORY_GRAPH tag is set to YES then doxygen will show the # dependencies a directory has on other directories in a graphical way. The # dependency relations are determined by the #include relations between the -# files in the directories. Explicit enabling a directory graph, when -# DIRECTORY_GRAPH is set to NO, can be accomplished by means of the command -# \directorygraph. Disabling a directory graph can be accomplished by means of -# the command \hidedirectorygraph. +# files in the directories. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. DIRECTORY_GRAPH = YES -# The DIR_GRAPH_MAX_DEPTH tag can be used to limit the maximum number of levels -# of child directories generated in directory dependency graphs by dot. -# Minimum value: 1, maximum value: 25, default value: 1. -# This tag requires that the tag DIRECTORY_GRAPH is set to YES. - -DIR_GRAPH_MAX_DEPTH = 1 - # The DOT_IMAGE_FORMAT tag can be used to set the image format of the images # generated by dot. For an explanation of the image formats see the section # output formats in the documentation of the dot tool (Graphviz (see: -# https://www.graphviz.org/)). +# http://www.graphviz.org/)). # Note: If you choose svg you need to set HTML_FILE_EXTENSION to xhtml in order # to make the SVG files visible in IE 9+ (other browsers do not have this # requirement). -# Possible values are: png, jpg, gif, svg, png:gd, png:gd:gd, png:cairo, +# Possible values are: png, png:cairo, png:cairo:cairo, png:cairo:gd, png:gd, +# png:gd:gd, jpg, jpg:cairo, jpg:cairo:gd, jpg:gd, jpg:gd:gd, gif, gif:cairo, +# gif:cairo:gd, gif:gd, gif:gd:gd, svg, png:gd, png:gd:gd, png:cairo, # png:cairo:gd, png:cairo:cairo, png:cairo:gdiplus, png:gdiplus and # png:gdiplus:gdiplus. # The default value is: png. @@ -2754,12 +2567,11 @@ DOT_PATH = DOTFILE_DIRS = -# You can include diagrams made with dia in doxygen documentation. Doxygen will -# then run dia to produce the diagram and insert it in the documentation. The -# DIA_PATH tag allows you to specify the directory where the dia binary resides. -# If left empty dia is assumed to be found in the default search path. +# The MSCFILE_DIRS tag can be used to specify one or more directories that +# contain msc files that are included in the documentation (see the \mscfile +# command). -DIA_PATH = +MSCFILE_DIRS = # The DIAFILE_DIRS tag can be used to specify one or more directories that # contain dia files that are included in the documentation (see the \diafile @@ -2768,10 +2580,10 @@ DIA_PATH = DIAFILE_DIRS = # When using plantuml, the PLANTUML_JAR_PATH tag should be used to specify the -# path where java can find the plantuml.jar file or to the filename of jar file -# to be used. If left blank, it is assumed PlantUML is not used or called during -# a preprocessing step. Doxygen will generate a warning when it encounters a -# \startuml command in this case and will not generate output for the diagram. +# path where java can find the plantuml.jar file. If left blank, it is assumed +# PlantUML is not used or called during a preprocessing step. Doxygen will +# generate a warning when it encounters a \startuml command in this case and +# will not generate output for the diagram. PLANTUML_JAR_PATH = @@ -2809,6 +2621,18 @@ DOT_GRAPH_MAX_NODES = 50 MAX_DOT_GRAPH_DEPTH = 0 +# Set the DOT_TRANSPARENT tag to YES to generate images with a transparent +# background. This is disabled by default, because dot on Windows does not seem +# to support this out of the box. +# +# Warning: Depending on the platform used, enabling this option may lead to +# badly anti-aliased labels on the edges of a graph (i.e. they become hard to +# read). +# The default value is: NO. +# This tag requires that the tag HAVE_DOT is set to YES. + +DOT_TRANSPARENT = NO + # Set the DOT_MULTI_TARGETS tag to YES to allow dot to generate multiple output # files in one run (i.e. multiple -o and -T options on the command line). This # makes dot run faster, but since only newer versions of dot (>1.8.10) support @@ -2821,8 +2645,6 @@ DOT_MULTI_TARGETS = NO # If the GENERATE_LEGEND tag is set to YES doxygen will generate a legend page # explaining the meaning of the various boxes and arrows in the dot generated # graphs. -# Note: This tag requires that UML_LOOK isn't set, i.e. the doxygen internal -# graphical representation for inheritance and collaboration diagrams is used. # The default value is: YES. # This tag requires that the tag HAVE_DOT is set to YES. @@ -2831,24 +2653,8 @@ GENERATE_LEGEND = YES # If the DOT_CLEANUP tag is set to YES, doxygen will remove the intermediate # files that are used to generate the various graphs. # -# Note: This setting is not only used for dot files but also for msc temporary -# files. +# Note: This setting is not only used for dot files but also for msc and +# plantuml temporary files. # The default value is: YES. DOT_CLEANUP = YES - -# You can define message sequence charts within doxygen comments using the \msc -# command. If the MSCGEN_TOOL tag is left empty (the default), then doxygen will -# use a built-in version of mscgen tool to produce the charts. Alternatively, -# the MSCGEN_TOOL tag can also specify the name an external tool. For instance, -# specifying prog as the value, doxygen will call the tool as prog -T -# -o . The external tool should support -# output file formats "png", "eps", "svg", and "ismap". - -MSCGEN_TOOL = - -# The MSCFILE_DIRS tag can be used to specify one or more directories that -# contain msc files that are included in the documentation (see the \mscfile -# command). - -MSCFILE_DIRS = diff --git a/docs/_static/kmm-logo-small.png b/docs/_static/kmm-logo-small.png deleted file mode 100644 index fb477477..00000000 Binary files a/docs/_static/kmm-logo-small.png and /dev/null differ diff --git a/docs/_static/kmm-logo.png b/docs/_static/kmm-logo.png index 83f08e61..a09bb2d3 100644 Binary files a/docs/_static/kmm-logo.png and b/docs/_static/kmm-logo.png differ diff --git a/examples/CMakeLists.txt b/examples/CMakeLists.txt index 782e7a22..988c8d91 100644 --- a/examples/CMakeLists.txt +++ b/examples/CMakeLists.txt @@ -1,115 +1,6 @@ -add_executable(vector_add_example ${PROJECT_SOURCE_DIR}/examples/vector_add.cu) -if(KMM_USE_CUDA) - set_target_properties( - vector_add_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) +# Every example is a .cu file with real GPU kernels, so it can only be compiled by nvcc/hipcc. +if(KMM_USE_CUDA OR KMM_USE_HIP) + add_subdirectory(hello_world) + add_subdirectory(vector_add) + add_subdirectory(reduce_sum) endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/vector_add.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(vector_add_example PRIVATE cxx_std_17) -target_link_libraries(vector_add_example PRIVATE kmm) - -add_executable(cpp_threads_vector_add_example ${PROJECT_SOURCE_DIR}/examples/cpp_threads_vector_add.cu) -if(KMM_USE_CUDA) - set_target_properties( - cpp_threads_vector_add_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/cpp_threads_vector_add.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(cpp_threads_vector_add_example PRIVATE cxx_std_17) -target_link_libraries(cpp_threads_vector_add_example PRIVATE kmm) - -add_executable(point_in_poly_example ${PROJECT_SOURCE_DIR}/examples/point_in_poly.cu) -if(KMM_USE_CUDA) - set_target_properties( - point_in_poly_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/point_in_poly.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(point_in_poly_example PRIVATE cxx_std_17) -target_link_libraries(point_in_poly_example PRIVATE kmm) - - -add_executable(matrix_multiply_example ${PROJECT_SOURCE_DIR}/examples/matrix_multiply.cu) -if(KMM_USE_CUDA) - set_target_properties( - matrix_multiply_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/matrix_multiply.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(matrix_multiply_example PRIVATE cxx_std_17) -target_link_libraries(matrix_multiply_example PRIVATE kmm) - - -add_executable(histogram_example ${PROJECT_SOURCE_DIR}/examples/histogram.cu) -if(KMM_USE_CUDA) - set_target_properties( - histogram_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/histogram.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(histogram_example PRIVATE cxx_std_17) -target_link_libraries(histogram_example PRIVATE kmm) - - -add_executable(reduction_example ${PROJECT_SOURCE_DIR}/examples/reduction.cu) -if(KMM_USE_CUDA) - set_target_properties( - reduction_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/reduction.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(reduction_example PRIVATE cxx_std_17) -target_link_libraries(reduction_example PRIVATE kmm) - - -add_executable(structs_example ${PROJECT_SOURCE_DIR}/examples/structs.cu) -if(KMM_USE_CUDA) - set_target_properties( - structs_example - PROPERTIES - CUDA_ARCHITECTURES "80" - CUDA_SEPARABLE_COMPILATION ON - CUDA_RESOLVE_DEVICE_SYMBOLS ON - ) -endif() -if(KMM_USE_HIP) - set_source_files_properties(${PROJECT_SOURCE_DIR}/examples/structs.cu PROPERTIES LANGUAGE HIP) -endif() -target_compile_features(structs_example PRIVATE cxx_std_17) -target_link_libraries(structs_example PRIVATE kmm) \ No newline at end of file diff --git a/examples/cpp_threads_vector_add.cu b/examples/cpp_threads_vector_add.cu deleted file mode 100644 index 4cca88d6..00000000 --- a/examples/cpp_threads_vector_add.cu +++ /dev/null @@ -1,107 +0,0 @@ -#include -#include - -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -__global__ void initialize_range(kmm::Range range, kmm::GPUSubviewMut output) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - if (i >= range.end) { - return; - } - - output[i] = float(i); -} - -__global__ void fill_range( - kmm::Range range, - float value, - kmm::GPUSubviewMut output -) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - if (i >= range.end) { - return; - } - - output[i] = value; -} - -__global__ void vector_add( - kmm::Range range, - kmm::GPUSubviewMut output, - kmm::GPUSubview left, - kmm::GPUSubview right -) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - - if (i >= range.end) { - return; - } - - output[i] = left[i] + right[i]; -} - -void main_loop(unsigned int id, kmm::RuntimeHandle& rt, long n, long chunk_size, dim3 block_size) { - using namespace kmm::placeholders; - auto A = kmm::Array {n}; - auto B = kmm::Array {n}; - auto C = kmm::Array {n}; - auto domain = kmm::TileDomain(n, chunk_size); - - rt.parallel_submit( // - domain, - kmm::GPUKernel(initialize_range, block_size), - _x, - write(A[_x]) - ); - - rt.parallel_submit( - domain, - kmm::GPUKernel(fill_range, block_size), - _x, - float(1.0), - write(B[_x]) - ); - - rt.parallel_submit( - domain, - kmm::GPUKernel(vector_add, block_size), - _x, - write(C[_x]), - A[_x], - B[_x] - ); - - auto result = std::vector(n); - C.copy_to(result); - - // Correctness check - for (long i = 0; i < n; i++) { - if (result[i] != float(i) + 1.0F) { - std::cerr << "[THREAD " << id << "] - wrong result at " << i << " : " << result[i] - << " != " << float(i) + 1 << std::endl; - return; - } - } -} - -int main() { - auto rt = kmm::make_runtime(); - spdlog::set_level(spdlog::level::warn); - long n = 200'000'000; - long chunk_size = n / 10; - dim3 block_size = 256; - unsigned int num_threads = 16; - std::vector threads; - - for (unsigned int thread = 0; thread < num_threads; thread++) { - threads.emplace_back(main_loop, thread, std::ref(rt), n, chunk_size, block_size); - } - for (unsigned int thread = 0; thread < num_threads; thread++) { - threads.at(thread).join(); - } - - std::cout << "Correctness check completed." << std::endl; - return EXIT_SUCCESS; -} diff --git a/examples/hello_world/CMakeLists.txt b/examples/hello_world/CMakeLists.txt new file mode 100644 index 00000000..61b0c9de --- /dev/null +++ b/examples/hello_world/CMakeLists.txt @@ -0,0 +1,18 @@ +add_executable(hello_world ${CMAKE_CURRENT_SOURCE_DIR}/main.cu) + +if(KMM_USE_CUDA) + set_target_properties( + hello_world + PROPERTIES + CUDA_ARCHITECTURES "80" + ) + # KMM_LAMBDA expands to a __device__-annotated lambda, which nvcc only accepts with this flag. + target_compile_options(hello_world PRIVATE $<$:--extended-lambda>) +endif() + +if(KMM_USE_HIP) + set_source_files_properties(${CMAKE_CURRENT_SOURCE_DIR}/main.cu PROPERTIES LANGUAGE HIP) +endif() + +target_compile_features(hello_world PRIVATE cxx_std_17) +target_link_libraries(hello_world PRIVATE kmm) diff --git a/examples/hello_world/main.cu b/examples/hello_world/main.cu new file mode 100644 index 00000000..692e4c20 --- /dev/null +++ b/examples/hello_world/main.cu @@ -0,0 +1,84 @@ +// A minimal example showing the round trip of a buffer through KMM: a vector is written on the +// CPU, doubled on the GPU, and read back on the CPU. +#if defined(KMM_USE_CUDA) + #include +#elif defined(KMM_USE_HIP) + #include +#endif +#include +#include + +#include "kmm/kmm.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/runtime.hpp" + +__global__ void double_values(kmm::ViewMut view) { + auto i = static_cast(blockIdx.x * blockDim.x + threadIdx.x); + + if (i < view.size()) { + view[i] *= 2.0f; + } +} + +int main() { + auto config = kmm::default_config_from_environment(); + config.host_memory_limit = 5ULL * 1024 * 1024 * 1024; + kmm::Context context = kmm::make_runtime(config); + auto host = context.host(); + auto dev = context.gpu(); + + // Write a vector on the CPU. + std::vector input = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f, 7.0f, 8.0f}; + auto array = host.from_vector(input); + + // Update the vector on the GPU: double every element. + dev.access(write(array)).submit([](auto stream, auto view) { + unsigned int block_size = 256; + unsigned int grid_size = + (static_cast(view.size()) + block_size - 1) / block_size; + + double_values<<>>(view); + }); + + unsigned int block_size = 256; + unsigned int grid_size = + (static_cast(array.size()) + block_size - 1) / block_size; + + host.prefetch(array, true); + + dev.submit( // + kmm::Kernel(double_values, grid_size, block_size), + write(array) + ); + + host.prefetch(array, true); + + dev.parallel_for( + array.shape(), + KMM_LAMBDA(auto index, auto view) { view[index] *= 2.0f; }, + write(array) + ); + + // Read the vector back on the CPU. + std::vector output(input.size()); + + { + auto guard = host.access(array); + const auto* ptr = guard.get().data(); + std::copy_n(ptr, input.size(), output.data()); + } + + std::cout << "Input: "; + for (float value : input) { + std::cout << value << ' '; + } + std::cout << '\n'; + + std::cout << "Output: "; + for (float value : output) { + std::cout << value << ' '; + } + std::cout << '\n'; + + return 0; +} diff --git a/examples/histogram.cu b/examples/histogram.cu deleted file mode 100644 index baa3d3dc..00000000 --- a/examples/histogram.cu +++ /dev/null @@ -1,93 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -void initialize_image( - unsigned long seed, - int width, - int height, - kmm::SubviewMut image -) { - std::mt19937 rand {seed}; - std::uniform_int_distribution dist {}; - - for (int i = 0; i < height; i++) { - for (int j = 0; j < width; j++) { - image[i][j] = dist(rand); - } - } -} - -void initialize_images( - kmm::Range subrange, - int width, - int height, - kmm::SubviewMut images -) { - for (auto i = subrange.begin; i < subrange.end; i++) { - initialize_image(i, width, height, images.drop_axis<0>(i)); - } -} - -__global__ void calculate_histogram( - kmm::Range image_ids, - int width, - int height, - kmm::GPUSubview images, - kmm::GPUSubviewMut histogram -) { - int image_id = blockIdx.z * blockDim.z + threadIdx.z + image_ids.begin; - int i = blockIdx.y * blockDim.y + threadIdx.y; - int j = blockIdx.x * blockDim.x + threadIdx.x; - - if (image_id < int(image_ids.end) && i < height && j < width) { - uint8_t value = images[image_id][i][j]; - atomicAdd(&histogram[image_id][value], 1); - } -} - -int main() { - using namespace kmm::placeholders; - spdlog::set_level(spdlog::level::trace); - - auto rt = kmm::make_runtime(); - int width = 1080; - int height = 1920; - int num_images = 2500; - int images_per_chunk = 500; - dim3 block_size = 256; - - auto histogram = kmm::Array {{256}}; - auto images = kmm::Array {{num_images, height, width}}; - - auto _imageid = kmm::Axis(2); - auto _i = kmm::Axis(1); - auto _j = kmm::Axis(0); - - rt.parallel_submit( - kmm::TileDomain(num_images, images_per_chunk), - kmm::Host(initialize_images), - _imageid, - width, - height, - write(images[_imageid][_][_]) - ); - - rt.synchronize(); - - rt.parallel_submit( - kmm::TileDomain({width, height, num_images}, {width, height, images_per_chunk}), - kmm::GPUKernel(calculate_histogram, block_size), - _imageid, - width, - height, - images[_imageid][_i][_j], - reduce(kmm::Reduction::Sum, privatize(_imageid), histogram[_]) - ); - - rt.synchronize(); - - return 0; -} diff --git a/examples/matrix_multiply.cu b/examples/matrix_multiply.cu deleted file mode 100644 index c74e1736..00000000 --- a/examples/matrix_multiply.cu +++ /dev/null @@ -1,136 +0,0 @@ -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -void fill_array(kmm::Bounds<2> region, kmm::SubviewMut array, float value) { - for (auto i = region.x.begin; i < region.x.end; i++) { - for (auto j = region.y.begin; j < region.y.end; j++) { - array[i][j] = value; - } - } -} - -void matrix_multiply( - kmm::DeviceResource& device, - kmm::Bounds<3> region, - int n, - int m, - int k, - kmm::GPUSubviewMut C, - kmm::GPUSubview A, - kmm::GPUSubview B -) { - using kmm::checked_cast; - - float alpha = 1.0; - float beta = 0.0; - - const float* A_ptr = A.data_at({region.y.begin, region.x.begin}); - const float* B_ptr = B.data_at({region.x.begin, region.z.begin}); - float* C_ptr = C.data_at({region.y.begin, region.z.begin}); - -#if __CUDA_ARCH__ - KMM_GPU_CHECK(cublasGemmEx( - device.blas(), - CUBLAS_OP_T, - CUBLAS_OP_T, - checked_cast(region.y.size()), - checked_cast(region.z.size()), - checked_cast(region.x.size()), - &alpha, - A_ptr, - CUDA_R_32F, - checked_cast(A.stride()), - B_ptr, - CUDA_R_32F, - checked_cast(B.stride()), - &beta, - C_ptr, - CUDA_R_32F, - checked_cast(C.stride()), - CUDA_R_32F, - CUBLAS_GEMM_DEFAULT - )); -#elif __HIP_DEVICE_COMPILE__ - KMM_GPU_CHECK(rocblas_gemm_ex( - device.blas(), - rocblas_operation_transpose, - rocblas_operation_transpose, - checked_cast(region.y.size()), - checked_cast(region.z.size()), - checked_cast(region.x.size()), - &alpha, - A_ptr, - rocblas_datatype_f32_r, - checked_cast(A.stride()), - B_ptr, - rocblas_datatype_f32_r, - checked_cast(B.stride()), - &beta, - C_ptr, - rocblas_datatype_f32_r, - checked_cast(C.stride()), - C_ptr, - rocblas_datatype_f32_r, - checked_cast(C.stride()), - rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, - 0, - 0 - )); -#endif -} - -int main() { - using namespace kmm::placeholders; - spdlog::set_level(spdlog::level::trace); - - auto rt = kmm::make_runtime(); - int n = 50000; - int m = 50000; - int k = 50000; - int chunk_size = n / 5; - - auto A = kmm::Array {{n, k}}; - auto B = kmm::Array {{k, m}}; - auto C = kmm::Array {{n, m}}; - - rt.parallel_submit( - {n, k}, - {chunk_size, chunk_size}, - kmm::Host(fill_array), - bounds(_x, _y), - write(A[_x][_y]), - 1.0F - ); - - rt.parallel_submit( - {k, m}, - {chunk_size, chunk_size}, - kmm::Host(fill_array), - bounds(_x, _y), - write(B[_x][_y]), - 1.0F - ); - - for (size_t repeat = 0; repeat < 1; repeat++) { - C.reset(); - - rt.parallel_submit( - {k, n, m}, - {chunk_size, chunk_size, chunk_size}, - kmm::GPU(matrix_multiply), - bounds(_x, _y, _z), - n, - m, - k, - reduce(kmm::Reduction::Sum, C[_y][_z]), - A[_y][_x], - B[_x][_z] - ); - - rt.synchronize(); - } - - return EXIT_SUCCESS; -} diff --git a/examples/point_in_poly.cu b/examples/point_in_poly.cu deleted file mode 100644 index fc89065b..00000000 --- a/examples/point_in_poly.cu +++ /dev/null @@ -1,112 +0,0 @@ -#ifdef KMM_USE_CUDA - #include -#elif KMM_USE_HIP - #include -#endif - -#include "spdlog/spdlog.h" - -#include "kmm/api/launcher.hpp" -#include "kmm/api/mapper.hpp" -#include "kmm/api/runtime_handle.hpp" - -__global__ void cn_pnpoly( - kmm::Range chunk, - kmm::GPUSubviewMut bitmap, - kmm::GPUSubview points, - int nvertices, - kmm::GPUView vertices -) { - int i = blockIdx.x * blockDim.x + threadIdx.x + chunk.begin; - - if (i < chunk.end) { - int c = 0; - float2 p = points[i]; - - int k = nvertices - 1; - - for (int j = 0; j < nvertices; k = j++) { // edge from v to vp - float2 vj = vertices[j]; - float2 vk = vertices[k]; - - float slope = (vk.x - vj.x) / (vk.y - vj.y); - - if (((vj.y > p.y) != (vk.y > p.y)) && //if p is between vj and vk vertically - (p.x < slope * (p.y - vj.y) + vj.x - )) { //if p.x crosses the line vj-vk when moved in positive x-direction - c = !c; - } - } - - bitmap[i] = c; // 0 if even (out), and 1 if odd (in) - } -} - -__global__ void init_points(kmm::Range chunk, kmm::GPUSubviewMut points) { - int i = blockIdx.x * blockDim.x + threadIdx.x + chunk.begin; - - if (i < chunk.end) { -#if __CUDA_ARCH__ - curandStatePhilox4_32_10_t state; - curand_init(1234, i, 0, &state); - points[i] = {curand_normal(&state), curand_normal(&state)}; -#elif __HIP_DEVICE_COMPILE__ - rocrand_state_philox4x32_10 state; - rocrand_init(1234, i, 0, &state); - points[i] = {rocrand_normal(&state), rocrand_normal(&state)}; -#endif - } -} - -void init_polygon(kmm::Range chunk, int nvertices, kmm::ViewMut vertices) { - for (int64_t i = chunk.begin; i < chunk.end; i++) { - float angle = float(i) / float(nvertices) * float(2.0F * M_PI); - vertices[i] = {cosf(angle), sinf(angle)}; - } -} - -int main() { - using namespace kmm::placeholders; - spdlog::set_level(spdlog::level::trace); - - auto rt = kmm::make_runtime(); - int nvertices = 1000; - int npoints = 1'000'000'000; - int npoints_per_chunk = npoints / 10; - dim3 block_size = 256; - - auto vertices = kmm::Array {nvertices}; - auto points = kmm::Array {npoints}; - auto bitmap = kmm::Array {npoints}; - - rt.submit( - kmm::ResourceId::host(), - kmm::Host(init_polygon), - kmm::Range(nvertices), - nvertices, - write(vertices) - ); - - rt.parallel_submit( - {npoints}, - {npoints_per_chunk}, - kmm::GPUKernel(init_points, block_size), - _x, - write(points[_x]) - ); - - rt.parallel_submit( - {npoints}, - {npoints_per_chunk}, - kmm::GPUKernel(cn_pnpoly, block_size), - _x, - write(bitmap[_x]), - points[_x], - nvertices, - vertices - ); - - rt.synchronize(); - - return EXIT_SUCCESS; -} diff --git a/examples/reduce_sum/CMakeLists.txt b/examples/reduce_sum/CMakeLists.txt new file mode 100644 index 00000000..2c61aeb4 --- /dev/null +++ b/examples/reduce_sum/CMakeLists.txt @@ -0,0 +1,16 @@ +add_executable(reduce_sum ${CMAKE_CURRENT_SOURCE_DIR}/main.cu) + +if(KMM_USE_CUDA) + set_target_properties( + reduce_sum + PROPERTIES + CUDA_ARCHITECTURES "80" + ) +endif() + +if(KMM_USE_HIP) + set_source_files_properties(${CMAKE_CURRENT_SOURCE_DIR}/main.cu PROPERTIES LANGUAGE HIP) +endif() + +target_compile_features(reduce_sum PRIVATE cxx_std_17) +target_link_libraries(reduce_sum PRIVATE kmm) diff --git a/examples/reduce_sum/main.cu b/examples/reduce_sum/main.cu new file mode 100644 index 00000000..f76aa485 --- /dev/null +++ b/examples/reduce_sum/main.cu @@ -0,0 +1,95 @@ +// An example showing how to use `kmm::Accumulator` to compute the sum of a large vector by +// splitting it across every available GPU: each device reduces its own slice into a single value, +// and KMM combines the per-device results into one final sum. +#include +#if defined(KMM_USE_CUDA) + #include +#elif defined(KMM_USE_HIP) + #include +#endif +#include +#include + +#include "kmm/kmm.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/runtime.hpp" + +static constexpr unsigned int BLOCK_SIZE = 256; +static constexpr unsigned int NUM_BLOCKS = 64; + +__global__ void sum_kernel(kmm::View input, kmm::ViewMut output) { + __shared__ float partial_sums[BLOCK_SIZE]; + + auto tid = threadIdx.x; + auto n = static_cast(input.size()); + auto grid_stride = gridDim.x * BLOCK_SIZE; + + float sum = 0.0f; + for (unsigned int i = blockIdx.x * BLOCK_SIZE + tid; i < n; i += grid_stride) { + sum += input[i]; + } + + partial_sums[tid] = sum; + __syncthreads(); + + for (unsigned int stride = BLOCK_SIZE / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + partial_sums[tid] += partial_sums[tid + stride]; + } + __syncthreads(); + } + + if (tid == 0) { + output[blockIdx.x] = partial_sums[0]; + } +} + +int main() { + auto config = kmm::default_config_from_environment(); + kmm::Context context = kmm::make_runtime(config); + + size_t n = 1'000'000; + std::vector input(n); + + for (size_t i = 0; i < n; i++) { + input[i] = 1.0f; + } + + auto values = context.from_vector(input); + + size_t num_devices = context.system_info().num_devices(); + size_t chunk_size = (n + num_devices - 1) / num_devices; + + auto accum = context.accumulator(kmm::ReductionOp::Sum); + + for (size_t d = 0; d < num_devices; d++) { + size_t begin = d * chunk_size; + size_t end = std::min(begin + chunk_size, n); + + if (begin >= end) { + break; + } + + auto chunk = values.slice_axis<0>(begin, end); + + context.gpu(kmm::DeviceId(d)) + .submit( + kmm::Kernel(sum_kernel, NUM_BLOCKS, BLOCK_SIZE), + chunk, + kmm::reduce(accum, NUM_BLOCKS) + ); + } + + auto output = accum.finalize(); + float result = context.to_scalar(output); + + std::cout << "Sum of " << n << " ones across " << num_devices << " device(s): " << result + << '\n'; + + if (result != static_cast(n)) { + std::cerr << "unexpected result!\n"; + return 1; + } + + return 0; +} diff --git a/examples/reduction.cu b/examples/reduction.cu deleted file mode 100644 index 2c128ee6..00000000 --- a/examples/reduction.cu +++ /dev/null @@ -1,167 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -__global__ void initialize_matrix_kernel( - kmm::Bounds<2, int> chunk, - kmm::GPUSubviewMut matrix -) { - int i = blockIdx.y * blockDim.y + threadIdx.y + chunk.y.begin; - int j = blockIdx.x * blockDim.x + threadIdx.x + chunk.x.begin; - - if (i < chunk.y.end && j < chunk.x.end) { - matrix[i][j] = float(i + 2 * j); - } -} - -__global__ void sum_total_kernel( - kmm::Bounds<2, int> chunk, - kmm::GPUSubview matrix, - kmm::GPUSubviewMut sum -) { - int i = blockIdx.y * blockDim.y + threadIdx.y + chunk.y.begin; - int j = blockIdx.x * blockDim.x + threadIdx.x + chunk.x.begin; - - if (i < chunk.y.end && j < chunk.x.end) { - sum[i][j] += matrix[i][j]; - } -} - -__global__ void sum_rows_kernel( - kmm::Bounds<2, int> chunk, - kmm::GPUSubview matrix, - kmm::GPUSubviewMut rows_sum -) { - int i = blockIdx.y * blockDim.y + threadIdx.y + chunk.y.begin; - int j = blockIdx.x * blockDim.x + threadIdx.x + chunk.x.begin; - - if (i < chunk.y.end && j < chunk.x.end) { - rows_sum[i][j] += matrix[i][j]; - } -} - -__global__ void sum_cols_kernel( - kmm::Bounds<2, int> chunk, - kmm::GPUSubview matrix, - kmm::GPUSubviewMut cols_sum -) { - int i = blockIdx.y * blockDim.y + threadIdx.y + chunk.y.begin; - int j = blockIdx.x * blockDim.x + threadIdx.x + chunk.x.begin; - - if (i < chunk.y.end && j < chunk.x.end) { - cols_sum[j][i] += matrix[i][j]; - } -} - -bool is_close(float expected, float gotten) { - return fabsf(expected - gotten) < fmaxf(1e-3F * fabsf(expected), 1e-9F); -} - -int run(kmm::RuntimeHandle& rt, int width, int height, int chunk_width, int chunk_height) { - using namespace kmm::placeholders; - auto domain = kmm::TileDomain({width, height}, {chunk_width, chunk_height}); - auto matrix = kmm::Array {{height, width}}; - - std::cout << "Execute for chunk size: " << chunk_width << "x" << chunk_height << "." - << std::endl; - - rt.parallel_submit( - domain, - kmm::GPUKernel(initialize_matrix_kernel, {16, 16}), - bounds(_x, _y), - write(matrix[_y][_x]) - ); - - rt.synchronize(); - - auto total_sum = kmm::Scalar(); - auto rows_sum = kmm::Array(height); - auto cols_sum = kmm::Array(width); - - rt.parallel_submit( - domain, - kmm::GPUKernel(sum_total_kernel, {16, 16}), - bounds(_x, _y), - matrix[_y][_x], - reduce(kmm::Reduction::Sum, privatize(_y, _x), total_sum) - ); - - rt.synchronize(); - - rt.parallel_submit( - domain, - kmm::GPUKernel(sum_rows_kernel, {16, 16}), - bounds(_x, _y), - matrix[_y][_x], - reduce(kmm::Reduction::Sum, privatize(_y), rows_sum[_x]) - ); - - rt.synchronize(); - - rt.parallel_submit( - domain, - kmm::GPUKernel(sum_cols_kernel, {16, 16}), - bounds(_x, _y), - matrix(_y, _x), - reduce(kmm::Reduction::Sum, privatize(_x), cols_sum[_y]) - ); - - rt.synchronize(); - - float total; - total_sum.copy_to(&total); - - if (!is_close(total, float(1.87125e+08))) { - std::cerr << "Wrong result for total_sum : " << total << " != " << float(1.87125e+08) - << std::endl; - return EXIT_FAILURE; - } - - std::vector rows; - rows_sum.copy_to(rows); - - for (int i = 0; i < height; i++) { - float expected = (float(width - 1) * 0.5F + float(2 * i)) * float(width); - - if (!is_close(rows[i], expected)) { - std::cerr << "Wrong result for rows_sum[" << i << "]: " << rows[i] << " != " << expected - << std::endl; - return EXIT_FAILURE; - } - } - - std::vector cols; - cols_sum.copy_to(cols); - - for (int i = 0; i < width; i++) { - float expected = (float(height - 1) + float(i)) * float(height); - - if (!is_close(cols[i], expected)) { - std::cerr << "Wrong result for cols_sum[" << i << "]: " << cols[i] << " != " << expected - << std::endl; - return EXIT_FAILURE; - } - } - - return EXIT_SUCCESS; -} - -int main() { - spdlog::set_level(spdlog::level::trace); - auto rt = kmm::make_runtime(); - int width = 500; - int height = 500; - - for (int nx = 1; nx <= 8; nx++) { - for (int ny = 1; ny <= 8; ny++) { - if (run(rt, width, height, width / nx, height / ny) != EXIT_SUCCESS) { - return EXIT_FAILURE; - } - } - } - - std::cout << "Correctness check completed." << std::endl; - return EXIT_SUCCESS; -} \ No newline at end of file diff --git a/examples/structs.cu b/examples/structs.cu deleted file mode 100644 index a82d520e..00000000 --- a/examples/structs.cu +++ /dev/null @@ -1,44 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -/// This defines the struct for the host-side code -struct Example { - int x; - kmm::Array y; -}; - -/// This defines the struct for the device-side code -struct ExampleView { - int x; - kmm::View y; -}; - -// This defines the fields of the `Example` struct -KMM_DEFINE_STRUCT_ARGUMENT(Example, it.x, it.y) - -// This defines that the "view" of `Example` is `ExampleView` -KMM_DEFINE_STRUCT_VIEW(Example, ExampleView) - -void example(kmm::Range range, ExampleView input) { - KMM_ASSERT(input.x == 123); - KMM_ASSERT(input.y.size() == 3); - KMM_ASSERT(input.y[0] == 1.0F); - KMM_ASSERT(input.y[1] == 2.0F); - KMM_ASSERT(input.y[2] == 3.0F); - std::cout << "input is correct for range " << range << "!" << std::endl; -} - -int main() { - using namespace kmm::placeholders; - auto rt = kmm::make_runtime(); - auto y = rt.allocate({1.0F, 2.0F, 3.0F}); - auto structure = Example {.x = 123, .y = y}; - - rt.parallel_submit(kmm::TileDomain(1000, 200), kmm::Host(example), _x, structure); - rt.synchronize(); - - return EXIT_SUCCESS; -} diff --git a/examples/vector_add.cu b/examples/vector_add.cu deleted file mode 100644 index cc31ff19..00000000 --- a/examples/vector_add.cu +++ /dev/null @@ -1,96 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/kmm.hpp" - -__global__ void initialize_range(kmm::Range range, kmm::GPUSubviewMut output) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - if (i >= range.end) { - return; - } - - output[i] = float(i); -} - -__global__ void fill_range( - kmm::Range range, - float value, - kmm::GPUSubviewMut output -) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - if (i >= range.end) { - return; - } - - output[i] = value; -} - -__global__ void vector_add( - kmm::Range range, - kmm::GPUSubviewMut output, - kmm::GPUSubview left, - kmm::GPUSubview right -) { - int64_t i = blockIdx.x * blockDim.x + threadIdx.x + range.begin; - - if (i >= range.end) { - return; - } - - output[i] = left[i] + right[i]; -} - -int main() { - using namespace kmm::placeholders; - - auto rt = kmm::make_runtime(); - spdlog::set_level(spdlog::level::trace); - long n = 200'000'000; - long chunk_size = n / 10; - dim3 block_size = 256; - - auto A = kmm::Array {n}; - auto B = kmm::Array {n}; - auto C = kmm::Array {n}; - auto domain = kmm::TileDomain(n, chunk_size); - - rt.parallel_submit( // - domain, - kmm::GPUKernel(initialize_range, block_size), - _x, - write(A[_x]) - ); - - rt.parallel_submit( - domain, - kmm::GPUKernel(fill_range, block_size), - _x, - float(1.0), - write(B[_x]) - ); - - rt.parallel_submit( - domain, - kmm::GPUKernel(vector_add, block_size), - _x, - write(C[_x]), - A[_x], - B[_x] - ); - - auto result = std::vector(n); - C.copy_to(result); - - // Correctness check - for (long i = 0; i < n; i++) { - if (result[i] != float(i) + 1.0F) { - std::cerr << "Wrong result at " << i << " : " << result[i] << " != " << float(i) + 1 - << std::endl; - return EXIT_FAILURE; - } - } - - std::cout << "Correctness check completed." << std::endl; - return EXIT_SUCCESS; -} diff --git a/examples/vector_add/CMakeLists.txt b/examples/vector_add/CMakeLists.txt new file mode 100644 index 00000000..aeb530b9 --- /dev/null +++ b/examples/vector_add/CMakeLists.txt @@ -0,0 +1,18 @@ +add_executable(vector_add ${CMAKE_CURRENT_SOURCE_DIR}/main.cu) + +if(KMM_USE_CUDA) + set_target_properties( + vector_add + PROPERTIES + CUDA_ARCHITECTURES "80" + ) + # KMM_LAMBDA expands to a __device__-annotated lambda, which nvcc only accepts with this flag. + target_compile_options(vector_add PRIVATE $<$:--extended-lambda>) +endif() + +if(KMM_USE_HIP) + set_source_files_properties(${CMAKE_CURRENT_SOURCE_DIR}/main.cu PROPERTIES LANGUAGE HIP) +endif() + +target_compile_features(vector_add PRIVATE cxx_std_17) +target_link_libraries(vector_add PRIVATE kmm) diff --git a/examples/vector_add/main.cu b/examples/vector_add/main.cu new file mode 100644 index 00000000..21e39923 --- /dev/null +++ b/examples/vector_add/main.cu @@ -0,0 +1,60 @@ +// A minimal example that adds two vectors on the GPU using `Device::parallel_for`: `c[i] = a[i] + +// b[i]` for every index `i`, with KMM taking care of allocating and moving the buffers. +#if defined(KMM_USE_CUDA) + #include +#elif defined(KMM_USE_HIP) + #include +#endif +#include +#include + +#include "kmm/kmm.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/runtime.hpp" + +int main() { + auto config = kmm::default_config_from_environment(); + kmm::Context context = kmm::make_runtime(config); + auto host = context.host(); + auto dev = context.gpu(); + + size_t n = 1'000'000; + std::vector input_a(n); + std::vector input_b(n); + + for (size_t i = 0; i < n; i++) { + input_a[i] = static_cast(i); + input_b[i] = static_cast(2 * i); + } + + // Move the input vectors to the runtime and allocate an (uninitialized) output vector. + auto a = host.from_vector(input_a); + auto b = host.from_vector(input_b); + auto c = context.empty(n); + + // Launch one GPU thread per element: `c[index] = a[index] + b[index]`. + dev.parallel_for( + c.shape(), + KMM_LAMBDA(auto index, auto va, auto vb, auto vc) { vc[index] = va[index] + vb[index]; }, + a, + b, + write(c) + ); + + std::vector output = context.to_vector(c); + + // Verify (and print) the first few results. + for (size_t i = 0; i < 10; i++) { + std::cout << input_a[i] << " + " << input_b[i] << " = " << output[i] << '\n'; + } + + for (size_t i = 0; i < n; i++) { + if (output[i] != input_a[i] + input_b[i]) { + std::cerr << "mismatch at index " << i << '\n'; + return 1; + } + } + + std::cout << "All " << n << " elements match!\n"; + return 0; +} diff --git a/external/unordered_dense b/external/unordered_dense new file mode 160000 index 00000000..e5b9441e --- /dev/null +++ b/external/unordered_dense @@ -0,0 +1 @@ +Subproject commit e5b9441ecf193f3e1e8f954527fc76edee20d7eb diff --git a/include/kmm/api/access.hpp b/include/kmm/api/access.hpp deleted file mode 100644 index aa1faeee..00000000 --- a/include/kmm/api/access.hpp +++ /dev/null @@ -1,130 +0,0 @@ -#pragma once - -#include "kmm/api/argument.hpp" -#include "kmm/api/mapper.hpp" - -namespace kmm { - -/** - * Encapsulates read access to an argument. - */ -template -struct Read { - Arg& argument; - M access_mapper = {}; -}; - -template -Read read(const Arg& argument, M access_mapper = {}) { - return {argument, access_mapper}; -} - -template -Read read(Read access) { - return access; -} - -/** - * Encapsulates write access to an argument. - */ -template -struct Write { - Arg& argument; - M access_mapper = {}; -}; - -template -Write write(Arg& argument, M access_mapper = {}) { - return {argument, access_mapper}; -} - -template -Write write(Read access) { - return {access.argument, access.access_mapper}; -} - -/** - * Encapsulates reduce access to an argument. - */ -template> -struct Reduce { - Arg& argument; - Reduction op; - M access_mapper = {}; - P private_mapper = {}; -}; - -template -struct Privatize { - M access_mapper; - - explicit Privatize(M access_mapper) : // - access_mapper(std::move(access_mapper)) {} -}; - -template -Privatize privatize(const M& mapper) { - return Privatize {mapper}; -} - -template -Privatize> privatize(const Is&... slices) { - return Privatize {bounds(slices...)}; -} - -template -Reduce reduce(Reduction op, Arg& argument, M access_mapper = {}) { - return {argument, op, access_mapper}; -} - -template -Reduce reduce( - Reduction op, - Privatize

private_mapper, - Arg& argument, - M access_mapper = {} -) { - return {argument, op, access_mapper, private_mapper.access_mapper}; -} - -template -Reduce reduce(Reduction op, Read access) { - return {access.argument, op, access.access_mapper}; -} - -template -Reduce reduce(Reduction op, Privatize

private_mapper, Read access) { - return {access.argument, op, access.access_mapper, private_mapper.access_mapper}; -} - -template typename Mode, size_t N, size_t I = 0> -struct MultiIndexAccess { - MultiIndexAccess(Arg& m_argument, MultiIndexMap m_mapper = {}) : - m_argument(m_argument), - m_mapper(m_mapper) {} - - template - auto operator[](const M& index) { - m_mapper.axes[I] = into_index_map(index); - - if constexpr (I + 1 == N) { - return Mode> {m_argument, {m_mapper}}; - } else { - return MultiIndexAccess(m_argument, m_mapper); - } - } - - private: - Arg& m_argument; - MultiIndexMap m_mapper = {}; -}; - -// Forward `Read` to `Read` only if `Arg` is not const -template -struct ArgumentHandler, std::enable_if_t>>: - ArgumentHandler> { - ArgumentHandler(Read access) : - ArgumentHandler>({access.argument, access.access_mapper}) {} -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/api/accumulator.hpp b/include/kmm/api/accumulator.hpp new file mode 100644 index 00000000..12bf1f86 --- /dev/null +++ b/include/kmm/api/accumulator.hpp @@ -0,0 +1,91 @@ +#pragma once + +#include + +#include "kmm/api/array.hpp" + +namespace kmm { + +/// A distributed reduction target: freshly allocated storage that kernels on one or more devices +/// add their contributions into. Each device writes its own partial result; the runtime combines +/// the per-device partials into the final array when `finalize()` is called. +template +class NDAccumulator { + KMM_NOT_COPYABLE(NDAccumulator) + + public: + using element_type = T; + using layout_type = LayoutT; + using domain_type = typename layout_type::domain_type; + using policy_type = typename layout_type::policy_type; + + /// Wraps a freshly allocated (uninitialized) `array` as a reduction target combining writes + /// using `op`. `array` must not have been written to yet. + explicit NDAccumulator(NDArray array, ReductionOp op = ReductionOp::Sum) : + m_array(std::move(array)), + m_op(op) { + m_array.buffer().runtime().begin_reduction(m_array.buffer().id(), data_type_of(), op); + } + + ~NDAccumulator() { + if (m_array.buffer()) { + m_array.buffer().runtime().rollback_reduction(m_array.buffer().id()); + } + } + + const layout_type& layout() const noexcept { + return m_array.layout(); + } + + ReductionOp op() const noexcept { + return m_op; + } + + const Buffer& buffer() const noexcept { + return m_array.buffer(); + } + + /// Combines everything written to this accumulator so far and returns the result. + /// Must be called at most once; after this, the `NDAccumulator` should be discarded. + NDArray finalize() { + KMM_ASSERT(m_array.buffer()); + m_array.buffer().runtime().finalize_reduction(m_array.buffer().id()); + return std::exchange(m_array, NDArray()); + } + + private: + NDArray m_array; + ReductionOp m_op; +}; + +template +class LaunchArg> { + public: + using resolve_type = NDView; + + explicit LaunchArg(const NDAccumulator& array) : m_array(array) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) { + m_index = requests.add(memory_id, m_array.buffer().id(), AccessMode::Reduce); + } + + resolve_type resolve(Runtime& runtime, const ResourceGrant& grant) { + auto accessor = grant.accessor(m_index); + auto* data = static_cast(accessor.address); + return resolve_type(data, m_array.layout()); + } + + void release(Runtime& runtime) {} + + private: + const NDAccumulator& m_array; + size_t m_index = 0; +}; + +template +using Accumulator = NDAccumulator, PolicyT>>; + +// `Reduce` / `ReduceInto` / `reduce(...)` -- the launch-argument side of reductions -- live in +// "kmm/api/reduce.hpp". + +} // namespace kmm diff --git a/include/kmm/api/argument.hpp b/include/kmm/api/argument.hpp deleted file mode 100644 index 18fd1d8a..00000000 --- a/include/kmm/api/argument.hpp +++ /dev/null @@ -1,92 +0,0 @@ -#pragma once - -#include "kmm/api/task_group.hpp" -#include "kmm/core/resource.hpp" - -namespace kmm { - -template -struct ArgumentHandler; - -template -struct ArgumentHandler: ArgumentHandler { - ArgumentHandler(const T& arg) : ArgumentHandler(arg) {} -}; - -template -struct ArgumentHandler: ArgumentHandler { - ArgumentHandler(T& arg) : ArgumentHandler(arg) {} -}; - -template -struct ArgumentHandler: ArgumentHandler { - ArgumentHandler(T&& arg) : ArgumentHandler(std::move(arg)) {} -}; - -template -using packed_argument_t = typename ArgumentHandler::type; - -template -packed_argument_t pack_argument(TaskInstance& task, T&& arg) { - return ArgumentHandler(std::forward(arg)).before_submit(task); -} - -template -struct ArgumentUnpack; - -template -auto unpack_argument(TaskContext& context, T&& arg) { - return ArgumentUnpack>::call(context, std::forward(arg)); -} - -template -struct Argument { - Argument(T value) : m_value(std::move(value)) {} - - static Argument pack(TaskInstance& builder, T value) { - return Argument {std::move(value)}; - } - - template - T unpack(TaskContext& context) { - return m_value; - } - - private: - T m_value; -}; - -template -struct ArgumentHandler { - using type = Argument; - - ArgumentHandler(T value) : m_value(std::move(value)) {} - - void initialize(const TaskGroupInit& init) { - // Nothing to do - } - - type before_submit(TaskInstance& builder) { - return Argument::pack(builder, m_value); - } - - void after_submit(const TaskSubmissionResult& result) { - // Nothing to do - } - - void commit(const TaskGroupCommit& commit) { - // Nothing to do - } - - private: - T m_value; -}; - -template -struct ArgumentUnpack> { - static auto call(TaskContext& context, Argument& data) { - return data.template unpack(context); - } -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/api/array.hpp b/include/kmm/api/array.hpp index de364f11..b5b3e146 100644 --- a/include/kmm/api/array.hpp +++ b/include/kmm/api/array.hpp @@ -1,496 +1,462 @@ #pragma once -#include -#include -#include +#include +#include +#include +#include +#include #include -#include "spdlog/spdlog.h" - -#include "kmm/api/access.hpp" -#include "kmm/api/argument.hpp" -#include "kmm/api/array_instance.hpp" -#include "kmm/api/view_argument.hpp" -#include "kmm/planner/read_planner.hpp" -#include "kmm/planner/reduction_planner.hpp" -#include "kmm/planner/write_planner.hpp" +#include "kmm/api/array_base.hpp" +#include "kmm/api/launch_arg.hpp" +#include "kmm/core/layout.hpp" +#include "kmm/core/view.hpp" +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memops/copy.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/resource.hpp" namespace kmm { -class ArrayBase { +template +class NDArray: public ArrayBase { public: - virtual ~ArrayBase() = default; - virtual const std::type_info& type_info() const = 0; - virtual size_t rank() const = 0; - virtual int64_t size(size_t axis) const = 0; - virtual const Runtime& runtime() const = 0; - virtual void synchronize() const = 0; - virtual void copy_bytes_to(void* output, size_t num_bytes) const = 0; -}; + using self_type = NDArray; + using element_type = T; + using layout_type = LayoutT; + static constexpr size_t rank = layout_type::rank; + using domain_type = typename layout_type::domain_type; + using policy_type = typename layout_type::policy_type; + using mapping_type = typename layout_type::mapping_type; + + using index_type = typename layout_type::index_type; + using ndindex_type = typename layout_type::ndindex_type; + using shape_type = typename layout_type::shape_type; + using range_type = typename layout_type::range_type; + using bounds_type = typename layout_type::bounds_type; + using stride_type = typename layout_type::stride_type; + using ndstrides_type = typename layout_type::ndstrides_type; + + /// The `NDArray` sharing this array's element type but over a different `Layout`, + template + using rebind_layout = NDArray; + + using zero_origin_type = rebind_layout; + using move_origin_type = rebind_layout; + using reverse_axes_type = rebind_layout; + + template + using permute_axes_type = + rebind_layout>; + + template + using swap_axes_type = rebind_layout>; + + using transpose_type = rebind_layout; + + template + using move_axis_to_position_type = + rebind_layout>; + + template + using move_axis_to_front_type = + rebind_layout>; + + template + using move_axis_to_back_type = + rebind_layout>; -template -class Array: public ArrayBase { - public: - Array(Dim shape = {}) : m_shape(shape) {} + template + using drop_axis_type = rebind_layout>; + + template + using insert_axis_type = rebind_layout>; - explicit Array(std::shared_ptr> b) : - m_instance(b), - m_shape(m_instance->distribution().array_size()) {} + template + using slice_axis_type = + rebind_layout>; - const std::type_info& type_info() const final { - return typeid(T); - } + template + using slice_type = rebind_layout>; - size_t rank() const final { - return N; - } + NDArray() = default; - Dim size() const { - return m_shape; - } + NDArray(layout_type layout) : m_layout(layout) {} - int64_t size(size_t axis) const final { - return m_shape.get_or_default(axis); - } + NDArray(domain_type domain, policy_type policy = {}) : + NDArray(make_layout(domain, policy).normalize_offset()) {} - int64_t volume() const { - return m_shape.volume(); - } + NDArray( + Runtime runtime, + layout_type layout, + std::optional fill_value = std::nullopt, + std::optional home = std::nullopt + ) : + ArrayBase( + Buffer::create( + runtime, + BufferLayout::for_type(static_cast(layout.offset_span().size())), + "", + fill_value ? FillValue::from(*fill_value) : FillValue {}, + home + ) + ), + m_layout(layout.normalize_offset()) {} - bool is_empty() const { - return m_shape.is_empty(); + NDArray( + Runtime runtime, + domain_type domain, + policy_type policy = {}, + std::optional fill_value = std::nullopt, + std::optional home = std::nullopt + ) : + NDArray(runtime, make_layout(domain, policy), fill_value, home) { + KMM_ASSERT(!domain.is_empty()); + KMM_ASSERT(!make_layout(domain, policy).is_empty()); + KMM_ASSERT(!make_layout(domain, policy).offset_span().is_empty()); } - bool has_instance() const { - return m_instance != nullptr; - } + /// Wrap a pre-existing, externally-owned allocation as an array, without copying. `external_ptr` + /// must point to at least `layout.offset_span().size()` elements of `T` resident in + /// `memory_id`, and the caller must keep it alive for as long as this array (or any copy) is in + /// use. KMM never allocates, frees, relocates, or copies the memory to another `MemoryId`. + NDArray(Runtime runtime, layout_type layout, T* external_ptr, MemoryId memory_id) : + ArrayBase( + Buffer::adopt( + runtime, + BufferLayout::for_type(static_cast(layout.offset_span().size())), + static_cast(external_ptr), + memory_id, + "" + ) + ), + m_layout(layout.normalize_offset()) {} - ArrayInstance& instance() const { - if (m_instance == nullptr) { - throw_uninitialized_array_exception(); - } + template + NDArray(const NDArray& that) : ArrayBase(that), m_layout(that.layout()) {} - return *m_instance; + const layout_type& layout() const noexcept { + return m_layout; } - const Distribution& distribution() const { - return instance().distribution(); + const domain_type& domain() const noexcept { + return m_layout.domain(); } - Dim chunk_size() const { - return distribution().chunk_size(); + const mapping_type& mapping() const noexcept { + return m_layout.mapping(); } - int64_t chunk_size(size_t axis) const { - return chunk_size().get_or_default(axis); + shape_type shape() const noexcept { + return m_layout.shape(); } - Runtime& runtime() const final { - return instance().runtime(); + /// The valid index range along the given axis. + range_type bounds(size_t axis) const noexcept { + return m_layout.bounds(axis); } - void synchronize() const final { - if (m_instance) { - m_instance->synchronize(); - } + /// The bounds (begin/end per axis) covered by this array. + bounds_type bounds() const noexcept { + return m_layout.bounds(); } - void reset() { - m_instance = nullptr; + /// The first valid index along the given axis. + index_type origin(size_t axis) const noexcept { + return m_layout.origin(axis); } - template - Read, M> access(M mapper = {}) { - return {*this, {std::move(mapper)}}; + /// The first valid index along each axis. + ndindex_type origin() const noexcept { + return m_layout.origin(); } - template - Read, M> access(M mapper = {}) const { - return {*this, {std::move(mapper)}}; + stride_type stride(size_t axis) const noexcept { + return m_layout.stride(axis); } - template - auto operator[](M first_index) { - return MultiIndexAccess, Read, N>(*this)[first_index]; + /// The stride along each axis. + ndstrides_type strides() const noexcept { + return m_layout.strides(); } - template - auto operator[](M first_index) const { - return MultiIndexAccess, Read, N>(*this)[first_index]; + index_type extent(size_t axis) const noexcept { + return m_layout.extent(axis); } - template - Read, MultiIndexMap> operator()(const Is&... index) { - return access(bounds(index...)); + /// The start of the valid range along the given axis. + index_type begin(size_t axis) const noexcept { + return m_layout.begin(axis); } - template - Read, MultiIndexMap> operator()(const Is&... index) const { - return access(bounds(index...)); + /// The end (exclusive) of the valid range along the given axis. + index_type end(size_t axis) const noexcept { + return m_layout.end(axis); } - void copy_bytes_to(void* output, size_t num_bytes) const { - KMM_ASSERT(num_bytes % sizeof(T) == 0); - KMM_ASSERT(is_equal(num_bytes / sizeof(T), volume())); - instance().copy_bytes_into(output); + /// The start of the valid range along each axis. + ndindex_type begin() const noexcept { + return m_layout.begin(); } - void copy_to(T* output) const { - instance().copy_bytes_into(output); + /// The end (exclusive) of the valid range along each axis. + ndindex_type end() const noexcept { + return m_layout.end(); } - template - void copy_to(T* output, I num_elements) const { - KMM_ASSERT(is_equal(num_elements, volume())); - instance().copy_bytes_into(output); + /// The total number of elements in this array. + index_type size() const noexcept { + return m_layout.size(); } - void copy_to(std::vector& output) const { - output.resize(checked_cast(volume())); - instance().copy_bytes_into(output.data()); + /// Whether this array covers zero elements. + bool is_empty() const noexcept { + return m_layout.is_empty(); } - std::vector copy_to_vector() const { - std::vector output(volume()); - copy_to(output); - return output; + template + DeviceEvent copy_to(NDArray& dst, MemoryId memory_id) const { + return runtime().submit_copy( + dst.buffer().id(), + buffer().id(), + make_copy_description(dst.layout(), m_layout, sizeof(T)), + memory_id + ); } - void copy_bytes_from(const void* input, size_t num_bytes) const { - KMM_ASSERT(num_bytes % sizeof(T) == 0); - KMM_ASSERT(is_equal(num_bytes / sizeof(T), volume())); - instance().copy_bytes_from(input); + template + DeviceEvent copy_to(NDArray& dst) const { + return copy_to(dst, dst.buffer().home().value_or(MemoryId::host())); } - void copy_from(T* input) const { - instance().copy_bytes_from(input); - } + template + rebind_layout> copy( + MemoryId home, + DstPolicyT dst_policy + ) const { + rebind_layout> + result(runtime(), domain(), dst_policy, std::nullopt, home); - template - void copy_from(T* input, I num_elements) const { - KMM_ASSERT(is_equal(num_elements, volume())); - instance().copy_bytes_from(input); + copy_to(result, home); + return result; } - void copy_from(const std::vector& input) const { - copy_from(input.data(), input.size()); + self_type copy(MemoryId home) const { + return copy(home, policy_type {}); } - private: - std::shared_ptr> m_instance; - Point m_offset; // Unused for now, always zero - Dim m_shape; -}; - -template -using Scalar = Array; - -template -struct ArgumentHandler>> { - using type = ViewArgument>; - - ArgumentHandler(Read> access) : - m_planner(access.argument.instance().shared_from_this()), - m_array_shape(access.argument.size()) {} - - void initialize(const TaskGroupInit& init) {} - - type before_submit(TaskInstance& task) { - auto region = Bounds(m_array_shape); - size_t buffer_index = task.add_buffer_requirement( // - m_planner.prepare_access(task.graph, task.memory_id, region, task.dependencies) - ); - - auto domain = views::dynamic_domain {region.size()}; - return {buffer_index, domain}; + self_type copy() const { + return copy(buffer().home().value_or(MemoryId::host())); } - void after_submit(const TaskSubmissionResult& result) { - m_planner.finalize_access(result.graph, result.event_id); + /// Returns this array rebased so its domain starts at the zero index. + zero_origin_type zero_origin() const noexcept { + return {buffer(), m_layout.zero_origin()}; } - void commit(const TaskGroupCommit& commit) { - m_planner.commit(commit.graph); + /// Returns this array shifted so it originates at the given index, keeping the same shape. + move_origin_type move_origin(ndindex_type new_origin) const noexcept { + return {buffer(), m_layout.move_origin(new_origin)}; } - private: - ArrayReadPlanner m_planner; - Dim m_array_shape; -}; - -template -struct ArgumentHandler>: ArgumentHandler>> { - ArgumentHandler(Array array) : ArgumentHandler>>(read(array)) {} -}; - -template -struct ArgumentHandler, M>> { - using type = ViewArgument>; - - static_assert( - is_dimensionality_accepted_by_mapper, - "mapper of 'read' must return N-dimensional region" - ); - - ArgumentHandler(Read, M> access) : - m_planner(access.argument.instance().shared_from_this()), - m_array_shape(access.argument.size()), - m_access_mapper(access.access_mapper) {} - - void initialize(const TaskGroupInit& init) {} - - type before_submit(TaskInstance& task) { - Bounds region = m_access_mapper(task.chunk, Bounds(m_array_shape)); - auto buffer_index = task.add_buffer_requirement( // - m_planner.prepare_access(task.graph, task.memory_id, region, task.dependencies) - ); - - auto domain = views::dynamic_subdomain {region.begin(), region.size()}; - return {buffer_index, domain}; + /// Returns this array restricted to the intersection of its bounds and the given bounds. + move_origin_type restrict_bounds(bounds_type new_bounds) const noexcept { + return {buffer(), m_layout.restrict_bounds(new_bounds)}; } - void after_submit(const TaskSubmissionResult& result) { - m_planner.finalize_access(result.graph, result.event_id); + /// Returns this array restricted along one axis to the intersection with [start, stop). + template + move_origin_type restrict_axis(index_type start, index_type stop) const noexcept { + return {buffer(), m_layout.template restrict_axis(start, stop)}; } - void commit(const TaskGroupCommit& commit) { - m_planner.commit(commit.graph); + /// Returns this array with the given axis dropped, fixed at the given index. + template + drop_axis_type drop_axis(index_type index) const noexcept { + return {buffer(), m_layout.template drop_axis(index)}; } - private: - ArrayReadPlanner m_planner; - Dim m_array_shape; - M m_access_mapper; -}; - -template -struct ArgumentHandler>> { - using type = ViewArgument>; - - ArgumentHandler(Write> access) : m_array(access.argument) {} - - void initialize(const TaskGroupInit& init) { - if (!m_array.has_instance()) { - auto instance = ArrayInstance::create( // - init.runtime, - map_domain_to_distribution(m_array.size(), init.domain, All()), - DataType::of() - ); - - m_array = Array(instance); - } - - m_planner = std::make_unique>(m_array.instance().shared_from_this()); + /// Returns this array with a new broadcast axis of the given extent inserted at the given + /// position. + template + insert_axis_type insert_axis( + index_type extent = static_cast(1) + ) const noexcept { + return {buffer(), m_layout.template insert_axis(extent)}; } - type before_submit(TaskInstance& task) { - auto access_region = Bounds(m_array.size()); - auto buffer_index = task.add_buffer_requirement( - m_planner->prepare_access(task.graph, task.memory_id, access_region, task.dependencies) - ); + /// Returns this array with the order of all axes reversed. + reverse_axes_type reverse_axes() const noexcept { + return {buffer(), m_layout.reverse_axes()}; + } - auto domain = views::dynamic_domain {access_region.size()}; - return {buffer_index, domain}; + /// Returns this array with its axes reordered according to the given permutation, e.g. + /// `permute_axes<2, 0, 1>()` moves the current axis 2 to position 0, axis 0 to position 1, + /// and axis 1 to position 2. + template + permute_axes_type permute_axes(IndexSequence seq = {}) const noexcept { + return {buffer(), m_layout.template permute_axes(seq)}; } - void after_submit(const TaskSubmissionResult& result) { - m_planner->finalize_access(result.graph, result.event_id); + /// Returns this array with axes `I` and `J` swapped. + template + swap_axes_type swap_axes() const noexcept { + return {buffer(), m_layout.template swap_axes()}; } - void commit(const TaskGroupCommit& commit) { - m_planner->commit(commit.graph); + /// Returns this array with axes 0 and 1 swapped. Only valid for a rank-2 array; use + /// `swap_axes` or `permute_axes` for other ranks. + transpose_type transpose() const noexcept { + return {buffer(), m_layout.transpose()}; } - private: - Array& m_array; - std::unique_ptr> m_planner; -}; + /// Returns this array with the given axis moved to the given position, preserving the + /// relative order of the remaining axes. + template + move_axis_to_position_type move_axis_to_position() const noexcept { + return {buffer(), m_layout.template move_axis_to_position()}; + } -template -struct ArgumentHandler, M>> { - using type = ViewArgument>; - - static_assert( - is_dimensionality_accepted_by_mapper, - "mapper of 'write' must return N-dimensional region" - ); - - ArgumentHandler(Write, M> access) : - m_array(access.argument), - m_shape(m_array.size()), - m_access_mapper(access.access_mapper) {} - - void initialize(const TaskGroupInit& init) { - if (!m_array.has_instance()) { - auto instance = ArrayInstance::create( // - init.runtime, - map_domain_to_distribution(m_array.size(), init.domain, m_access_mapper), - DataType::of() - ); - - m_array = Array(instance); - } + /// Returns this array with the given axis moved to the front (position 0), preserving the + /// relative order of the remaining axes. + template + move_axis_to_front_type move_axis_to_front() const noexcept { + return {buffer(), m_layout.template move_axis_to_front()}; + } - m_planner = std::make_unique>(m_array.instance().shared_from_this()); + /// Returns this array with the given axis moved to the back (position `rank - 1`), + /// preserving the relative order of the remaining axes. + template + move_axis_to_back_type move_axis_to_back() const noexcept { + return {buffer(), m_layout.template move_axis_to_back()}; } - type before_submit(TaskInstance& task) { - auto access_region = m_access_mapper(task.chunk, Bounds(m_shape)); - auto buffer_index = task.add_buffer_requirement( - m_planner->prepare_access(task.graph, task.memory_id, access_region, task.dependencies) - ); + /// Returns this array with the given axis sliced according to the given slice token (e.g. + /// `all`, a `Range`, `new_axis`). + template + slice_axis_type slice_axis(const SliceT& slice) const noexcept { + return {buffer(), m_layout.template slice_axis(slice)}; + } - auto domain = views::dynamic_subdomain {access_region.begin(), access_region.size()}; - return {buffer_index, domain}; + /// Returns this array with the given axis narrowed to the range [start, end). + template + self_type slice_axis(index_type start, index_type end) const noexcept { + return self_type(buffer(), m_layout.template slice_axis(start, end)); } - void after_submit(const TaskSubmissionResult& result) { - m_planner->finalize_access(result.graph, result.event_id); + /// Returns this array sliced across all axes at once, one slice token per axis. + template + slice_type slice(const Slices&... slices) const noexcept { + return {buffer(), m_layout.slice(slices...)}; } - void commit(const TaskGroupCommit& commit) { - m_planner->commit(commit.graph); + template + slice_axis_type<0, SliceT> operator[](const SliceT& slice) const noexcept { + return {buffer(), m_layout.template slice_axis<0>(slice)}; } private: - Array& m_array; - Dim m_shape; - M m_access_mapper; - std::unique_ptr> m_planner; + template + friend class NDArray; + + NDArray(Buffer buffer, layout_type layout) noexcept : + ArrayBase(std::move(buffer)), + m_layout(layout) {} + + layout_type m_layout {}; }; -template -struct ArgumentHandler>> { - using type = ViewArgument>; - - ArgumentHandler(Reduce> access) : - m_array(access.argument), - m_operation(access.op) {} - - void initialize(const TaskGroupInit& init) { - if (!m_array.has_instance()) { - auto instance = ArrayInstance::create( // - init.runtime, - map_domain_to_distribution( // - m_array.size(), - init.domain, - All(), - true - ), - DataType::of() - ); - - m_array = Array(instance); - } +template +using Array = NDArray, PolicyT>>; - m_planner = std::make_unique>( - m_array.instance().shared_from_this(), - m_operation - ); - } +template +using SubArray = NDArray, PolicyT>>; - type before_submit(TaskInstance& task) { - auto access_region = Bounds(m_array.size()); +template +using StridedArray = Array; - size_t buffer_index = task.add_buffer_requirement( - m_planner - ->prepare_access(task.graph, task.memory_id, access_region, 1, task.dependencies) - ); +template +using StridedSubArray = SubArray; - views::dynamic_domain domain = {access_region.size()}; +template +using Scalar = NDArray, RowMajor>>; - return {buffer_index, domain}; - } +namespace detail { - void after_submit(const TaskSubmissionResult& result) { - m_planner->finalize_access(result.graph, result.event_id); +/// Shared implementation of `LaunchArg` for `NDArray`, parameterized on the access +/// mode granted to the resolved view: `AccessMode::Read` yields a `NDView`, +/// `AccessMode::ReadWrite` a `NDView`. +template +class LaunchArgArray { + public: + using view_element_type = std::conditional_t; + using resolve_type = NDView; + + explicit LaunchArgArray(const NDArray& array) : m_array(array) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) { + m_index = requests.add(memory_id, m_array.buffer().id(), Mode); } - void commit(const TaskGroupCommit& commit) { - m_planner->commit(commit.graph); + resolve_type resolve(Runtime& runtime, const ResourceGrant& grant) { + auto accessor = grant.accessor(m_index); + auto* data = static_cast(accessor.address); + return resolve_type(data, m_array.layout()); } + void release(Runtime& runtime) {} + private: - Array& m_array; - Reduction m_operation; - std::unique_ptr> m_planner; + NDArray m_array; + size_t m_index = 0; }; -template -struct ArgumentHandler, M, P>> { - static constexpr size_t K = mapper_dimensionality

; - using type = ViewArgument>; - - static_assert( - is_dimensionality_accepted_by_mapper, - "mapper of 'reduce' must return N-dimensional region" - ); - - static_assert( - is_dimensionality_accepted_by_mapper, - "private mapper of 'reduce' must return K-dimensional region" - ); - - ArgumentHandler(Reduce, M, P> access) : - m_array(access.argument), - m_operation(access.op), - m_access_mapper(access.access_mapper), - m_private_mapper(access.private_mapper) {} - - void initialize(const TaskGroupInit& init) { - if (!m_array.has_instance()) { - auto instance = ArrayInstance::create( // - init.runtime, - map_domain_to_distribution( // - m_array.size(), - init.domain, - m_access_mapper, - true - ), - DataType::of() - ); - - m_array = Array(instance); - } - - m_planner = std::make_unique>( - m_array.instance().shared_from_this(), - m_operation - ); - } - - type before_submit(TaskInstance& task) { - auto access_region = m_access_mapper(task.chunk, Bounds(m_array.size())); - auto private_region = m_private_mapper(task.chunk); - - auto rep = checked_cast(private_region.volume()); - size_t buffer_index = task.add_buffer_requirement( - m_planner - ->prepare_access(task.graph, task.memory_id, access_region, rep, task.dependencies) - ); +} // namespace detail - views::dynamic_subdomain domain = { - concat(private_region, access_region).begin(), - concat(private_region, access_region).size() - }; +/// Read-only access to a `NDArray` (the default when passed to `Device::scope`/`Host::scope` unwrapped). +template +class LaunchArg>: public detail::LaunchArgArray { + public: + using detail::LaunchArgArray::LaunchArgArray; +}; - return {buffer_index, domain}; - } +/// Read-only access to a `NDArray` explicitly wrapped in `read(...)`. +template +class LaunchArg>>: + public detail::LaunchArgArray { + public: + explicit LaunchArg(Read> arg) : + detail::LaunchArgArray(arg.value) {} +}; - void after_submit(const TaskSubmissionResult& result) { - m_planner->finalize_access(result.graph, result.event_id); - } +/// Read-write access to a `NDArray` wrapped in `write(...)`. If `arg.value` is not yet +/// allocated (i.e. it has no backing buffer, only a layout), a new array is allocated on `runtime` +/// during `acquire()` and assigned back to `arg.value`, so the caller observes the allocation too. +template +class LaunchArg>>: + public detail::LaunchArgArray { + public: + explicit LaunchArg(Write> arg) : + detail::LaunchArgArray(arg.value), + m_target(&arg.value) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) { + if (!*m_target) { + *m_target = NDArray(runtime, m_target->layout()); + *this = LaunchArg(write(*m_target)); + } - void commit(const TaskGroupCommit& commit) { - m_planner->commit(commit.graph); + detail::LaunchArgArray::acquire( + runtime, + requests, + memory_id + ); } private: - Array& m_array; - Reduction m_operation; - std::unique_ptr> m_planner; - M m_access_mapper; - P m_private_mapper; + NDArray* m_target; }; -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/api/array_base.hpp b/include/kmm/api/array_base.hpp new file mode 100644 index 00000000..276b3296 --- /dev/null +++ b/include/kmm/api/array_base.hpp @@ -0,0 +1,33 @@ +#pragma once + +#include + +#include "kmm/api/buffer.hpp" + +namespace kmm { + +class ArrayBase { + public: + ArrayBase() = default; + ArrayBase(const ArrayBase&) = default; + + const Buffer& buffer() const { + return m_buffer; + } + + Runtime runtime() const noexcept { + return m_buffer.runtime(); + } + + explicit operator bool() const { + return bool(m_buffer); + } + + protected: + explicit ArrayBase(Buffer buffer) noexcept : m_buffer(std::move(buffer)) {} + + private: + Buffer m_buffer; +}; + +} // namespace kmm diff --git a/include/kmm/api/array_instance.hpp b/include/kmm/api/array_instance.hpp deleted file mode 100644 index 0dd4091a..00000000 --- a/include/kmm/api/array_instance.hpp +++ /dev/null @@ -1,35 +0,0 @@ -#pragma once - -#include "kmm/planner/array_descriptor.hpp" - -namespace kmm { - -class Runtime; - -template -class ArrayInstance: - public ArrayDescriptor, - public std::enable_shared_from_this> { - KMM_NOT_COPYABLE_OR_MOVABLE(ArrayInstance) - - ArrayInstance(TaskGraph& stage, Runtime& rt, Distribution dist, DataType dtype); - - public: - static std::shared_ptr create(Runtime& rt, Distribution dist, DataType dtype); - ~ArrayInstance(); - - void copy_bytes_into(void* data); - void copy_bytes_from(const void* data); - void synchronize() const; - - Runtime& runtime() const { - return *m_rt; - } - - private: - std::shared_ptr m_rt; -}; - -[[noreturn]] void throw_uninitialized_array_exception(); - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/api/buffer.hpp b/include/kmm/api/buffer.hpp new file mode 100644 index 00000000..94fb1737 --- /dev/null +++ b/include/kmm/api/buffer.hpp @@ -0,0 +1,91 @@ +#pragma once + +#include +#include + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memops/copy.hpp" +#include "kmm/runtime/runtime.hpp" +#include "kmm/utils/refcnt_ptr.hpp" + +namespace kmm { + +class Runtime; + +class Buffer { + public: + struct Impl; + + Buffer() = default; + + /// Create a new runtime-managed buffer. The runtime owns the allocation and is free to place, + /// move, and evict it across memories as needed. + static Buffer create( + Runtime runtime, + BufferLayout layout, + std::string name, + FillValue fill_value = {}, + std::optional home = {}, + std::optional kind = std::nullopt + ); + + /// Wrap a pre-existing, externally-owned allocation as a buffer, without copying. KMM never + /// allocates, frees, or relocates the memory; the caller must keep `ptr` alive for as long as + /// the returned `Buffer` (or any copy of it) is in use. See `Runtime::adopt_buffer`. + static Buffer adopt( + Runtime runtime, + BufferLayout layout, + void* ptr, + MemoryId memory_id, + std::string name = {} + ); + + BufferId id() const; + BufferLayout layout() const; + Runtime runtime() const; + + /// Returns this buffer's home memory, if it has one. See `Runtime::buffer_home`. + std::optional home() const; + + void prefetch(MemoryId memory_id, bool invalidate_others = false) const; + void poison(std::exception_ptr reason) const; + void invalidate() const; + + void copy_to( + void* dest, + size_t nbytes, + size_t offset = 0, + MemoryId memory_id = MemoryId::host() + ) const; + + void copy_from( + const void* dest, + size_t nbytes, + size_t offset = 0, + MemoryId memory_id = MemoryId::host() + ) const; + + void copy_to( // + void* dest, + CopyDescription description, + MemoryId memory_id = MemoryId::host() + ) const; + + void copy_from( + const void* src, + CopyDescription description, + MemoryId memory_id = MemoryId::host() + ) const; + + explicit operator bool() const { + return bool(m_impl); + } + + private: + refcnt_ptr m_impl; +}; + +KMM_REFCNT_TRAITS_FWD(Buffer::Impl) + +} // namespace kmm diff --git a/include/kmm/api/context.hpp b/include/kmm/api/context.hpp new file mode 100644 index 00000000..e3c1cbc8 --- /dev/null +++ b/include/kmm/api/context.hpp @@ -0,0 +1,271 @@ +#pragma once + +#include +#include +#include +#include + +#include "kmm/api/accumulator.hpp" +#include "kmm/api/array.hpp" +#include "kmm/api/reduce.hpp" +#include "kmm/api/resource_guard.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/memory_manager.hpp" +#include "kmm/runtime/resource.hpp" +#include "kmm/runtime/runtime.hpp" + +namespace kmm { + +class Host; +class Device; + +class Context { + public: + Context(Runtime runtime, MemoryTransaction transaction = {}) : + m_runtime(std::move(runtime)), + m_transaction(std::move(transaction)) {} + + template + NDArray> array( + DomainT domain, + PolicyT policy = {}, + std::optional fill_value = {} + ) { + return NDArray>( + runtime(), + domain, + policy, + fill_value, + affinity_memory_id() + ); + } + + template + NDArray> adopt( + T* data, + DomainT domain, + std::optional memory_id = {}, + PolicyT policy = {} + ) { + return NDArray>( + runtime(), + make_layout(domain, policy), + data, + memory_id.value_or(affinity_memory_id()) + ); + } + + template + Array empty(ExtentT... extents) { + return array(shape(extents...), PolicyT {}); + } + + template + Array fill(T value, ExtentT... extents) { + return array(shape(extents...), PolicyT {}, value); + } + + template + Array ones(ExtentT... extents) { + return array(shape(extents...), PolicyT {}, static_cast(1)); + } + + template + Array zeros(ExtentT... extents) { + return array(shape(extents...), PolicyT {}, T {}); + } + + template + NDArray empty_like(const NDArray& that) { + return array(that.domain(), typename LayoutT::policy_type {}); + } + + template + NDArray ones_like(const NDArray& that) { + return array(that.domain(), typename LayoutT::policy_type {}, static_cast(1)); + } + + template + NDArray zeros_like(const NDArray& that) { + return array(that.domain(), typename LayoutT::policy_type {}, T {}); + } + + template + Array from_vector(const T* data, size_t nelem) { + auto result = empty(nelem); + result.buffer().copy_from(data, nelem * sizeof(T), 0, affinity_memory_id()); + return result; + } + + template + Array from_vector(const std::vector& input) { + return from_vector(input.data(), input.size()); + } + + template + std::vector to_vector(Array input) { + size_t offset = checked_cast(input.layout().base_offset()); + size_t nelem = checked_cast(input.layout().size()); + + std::vector result; + result.resize(nelem); + input.buffer().copy_to( + result.data(), + nelem * sizeof(T), + offset * sizeof(T), + affinity_memory_id() + ); + return result; + } + + template + Array from_scalar(const T& value) { + return from_vector(&value, 1); + } + + template + T to_scalar(const NDArray& input) { + KMM_ASSERT(input.layout().size() == 1); + size_t offset = checked_cast(input.layout().base_offset()); + + T result; + input.buffer().copy_to(&result, sizeof(T), offset * sizeof(T), affinity_memory_id()); + return result; + } + + template + NDArray copy(const NDArray& that) { + auto result = empty_like(that); + copy(result, that); + return result; + } + + template + DeviceEvent copy(NDArray& dst, const NDArray& src) { + return copy(dst, src, affinity_memory_id()); + } + + template + DeviceEvent copy( + NDArray& dst, + const NDArray& src, + MemoryId dst_memory_id + ) { + return m_runtime.submit_copy( + dst.buffer().id(), + src.buffer().id(), + make_copy_description(dst.layout(), src.layout(), sizeof(T)), + dst_memory_id, + std::nullopt, + transaction() + ); + } + + template + DeviceEvent reduce_into( + ReductionOp op, + const NDArray& src, + const NDArray& dst + ) { + auto desc = + make_reduction_description(dst.layout(), src.layout(), 0, data_type_of(), op); + + return m_runtime.submit_reduction( + dst.buffer().id(), + src.buffer().id(), + desc, + affinity_memory_id(), + affinity_stream(), + m_transaction + ); + } + + template + T reduce(ReductionOp op, const NDArray& src) { + Scalar result; + reduce_into(op, src, result); + return to_scalar(result); + } + + template + T sum(const NDArray& src) { + return reduce(ReductionOp::Sum, src); + } + + void prefetch(const Buffer& buffer, bool invalidate_others = false) { + buffer.prefetch(affinity_memory_id(), invalidate_others); + } + + template + void prefetch(const NDArray& src, bool invalidate_others = false) { + prefetch(src.buffer(), invalidate_others); + } + + template + NDAccumulator, PolicyT>> accumulator( + ReductionOp op, + ExtentT... extents + ) { + return NDAccumulator, PolicyT>>( + empty(extents...), + op + ); + } + + template + NDAccumulator, PolicyT>> sum_accumulator( + ExtentT... extents + ) { + return accumulator(ReductionOp::Sum, extents...); + } + + const Runtime& runtime() const noexcept { + return m_runtime; + } + + Runtime& runtime() noexcept { + return m_runtime; + } + + const MemoryTransaction& transaction() const noexcept { + return m_transaction; + } + + const SystemInfo& system_info() const noexcept { + return m_runtime.system_info(); + } + + void synchronize() { + m_runtime.synchronize(); + } + + void synchronize(const DeviceEvent& event) { + m_runtime.synchronize(event); + } + + void synchronize(const DeviceEventSet& events) { + m_runtime.synchronize(events); + } + + /// Returns a `Host` context sharing this context's runtime and transaction. + Host host(); + + /// Returns a `Device` context targeting `device_id`, sharing this context's runtime and + /// transaction. + Device gpu(DeviceId device_id = DeviceId(0)); + + virtual MemoryId affinity_memory_id() const noexcept { + return MemoryId::host(); + } + + virtual std::optional affinity_stream() const noexcept { + return std::nullopt; + } + + protected: + Runtime m_runtime; + MemoryTransaction m_transaction; +}; + +} // namespace kmm diff --git a/include/kmm/api/device.hpp b/include/kmm/api/device.hpp new file mode 100644 index 00000000..fab46aad --- /dev/null +++ b/include/kmm/api/device.hpp @@ -0,0 +1,148 @@ +#pragma once + +#include +#include +#include + +#include "kmm/api/context.hpp" +#include "kmm/api/parallel_for.hpp" +#include "kmm/api/parallel_reduce.hpp" +#include "kmm/api/resource_guard.hpp" +#include "kmm/core/shape.hpp" +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/device_stream.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/utils/gpu_utils.hpp" +#include "kmm/utils/refcnt_ptr.hpp" + +namespace kmm { + +template +class DeviceGuard; + +class Device: public Context { + public: + explicit Device(Runtime runtime, DeviceId device_id, MemoryTransaction transaction = {}); + + DeviceStream stream() const noexcept; + + DeviceId device_id() const noexcept { + return m_device_id; + } + + operator g_stream_t() const noexcept { + return *m_stream; + } + + /// Acquires `args` and calls `fun(stream, views...)`, where `stream` is this device's + /// `GPUStream` and `views...` are the resolved views for `args...`. + template + decltype(auto) submit(F&& fun, Args&&... args) { + return access(std::forward(args)...).submit(std::forward(fun)); + } + + template + DeviceGuard access(Args&&... args) { + return DeviceGuard( + runtime(), + m_device_id, + m_stream, + transaction(), + std::forward(args)... + ); + } + + template + void parallel_for(Shape shape, F fun, Args&&... args) { + submit(ParallelFor(shape, std::move(fun)), args...); + } + + template< + size_t N, + typename F, + typename... Args, + typename OutputT = + std::invoke_result_t, typename LaunchArg::resolve_type...>> + OutputT parallel_reduce(Shape shape, F fun, Args&&... args) { + auto redux = ParallelReduce(shape, fun); + auto partials = empty(redux.num_outputs()); + submit(redux, write(partials), args...); + return this->sum(partials); + } + + MemoryId affinity_memory_id() const noexcept override { + return MemoryId::device(m_device_id); + } + + private: + DeviceId m_device_id; + std::shared_ptr m_stream; +}; + +template +class DeviceGuard { + public: + DeviceGuard( + Runtime& runtime, + DeviceId device_id, + std::shared_ptr stream, + MemoryTransaction parent, + Args... args + ) : + m_device_id(device_id), + m_stream(stream), + m_guard(runtime, MemoryId::device(device_id), *stream, parent, args...) { + auto& registry = runtime.event_registry(); + m_stream_id = registry.lookup_or_register_stream(*m_stream); + registry.wait_on_event(m_stream_id, m_guard.dependencies()); + } + + ~DeviceGuard() { + auto event = m_guard.runtime().event_registry().record(m_stream_id); + m_guard.release(event); + } + + Device context() const { + return Device(m_guard.runtime(), m_device_id, m_guard.transaction()); + } + + DeviceId device_id() const noexcept { + return m_device_id; + } + + DeviceStreamId stream_id() const noexcept { + return m_stream_id; + } + + GPUStreamRef stream() const noexcept { + return *m_stream; + } + + template + decltype(auto) get() { + return m_guard.template get(); + } + + void poison(std::exception_ptr reason) noexcept { + m_guard.poison(std::move(reason)); + } + + template + void submit(F fun) { + m_guard.apply(std::move(fun), context()); + } + + template + void operator>>(F fun) { + submit(std::move(fun)); + } + + private: + DeviceId m_device_id; + DeviceStreamId m_stream_id; + std::shared_ptr m_stream; + ResourceGuard m_guard; +}; + +} // namespace kmm diff --git a/include/kmm/api/dist_array.hpp b/include/kmm/api/dist_array.hpp new file mode 100644 index 00000000..cb66684b --- /dev/null +++ b/include/kmm/api/dist_array.hpp @@ -0,0 +1,144 @@ +#pragma once + +#include +#include +#include +#include + +#include "kmm/api/array.hpp" +#include "kmm/api/context.hpp" +#include "kmm/api/device.hpp" +#include "kmm/api/parallel_for.hpp" +#include "kmm/core/bounds.hpp" +#include "kmm/core/distribution.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/identifiers.hpp" + +namespace kmm { + +/// A large N-dimensional array stored as a set of smaller chunks, each chunk being an +/// independent `Array` (own buffer) with its own home `MemoryId`. +template +class DistArray { + public: + using element_type = T; + using policy_type = PolicyT; + static constexpr size_t rank = N; + using shape_type = Shape; + using point_type = Point; + using bounds_type = Bounds; + using chunk_type = Array; + using slice_type = typename chunk_type::move_origin_type; + + DistArray() = default; + + /// Allocates one chunk (its own `Buffer`) per grid cell of `dist`. `home_fn(linear_index)` + /// picks the home `MemoryId` of the chunk with the given row-major linear chunk index; if + /// omitted, chunks are round-robined over the available GPUs + /// (`runtime.system_info().num_devices()`), falling back to the host if there are none. + DistArray( + Runtime runtime, + Distribution dist, + PolicyT policy = {}, + std::optional fill_value = std::nullopt, + std::function home_fn = {} + ) : + m_dist(dist) { + if (!home_fn) { + size_t num_devices = runtime.system_info().num_devices(); + home_fn = [num_devices](size_t linear_index) { + return num_devices > 0 ? MemoryId::device(DeviceId(linear_index % num_devices)) + : MemoryId::host(); + }; + } + + m_chunks.reserve(dist.num_chunks()); + + for (size_t i = 0; i < dist.num_chunks(); i++) { + auto extent = dist.chunk_extent(dist.unravel(i)); + auto home = home_fn(i); + + m_chunks.push_back( + ChunkEntry {chunk_type(runtime, extent, policy, fill_value, home), home} + ); + } + } + + /// The partitioning geometry (total shape, chunk shape, grid shape) of this array. + const Distribution& distribution() const noexcept { + return m_dist; + } + + /// The full (global) extent of this array. + shape_type shape() const noexcept { + return m_dist.total_shape(); + } + + /// The number of chunks along each axis. + shape_type grid_shape() const noexcept { + return m_dist.grid_shape(); + } + + /// The total number of chunks. + size_t num_chunks() const noexcept { + return m_chunks.size(); + } + + const chunk_type& chunk(size_t linear_index) const noexcept { + return m_chunks[linear_index].array; + } + + chunk_type& chunk(size_t linear_index) noexcept { + return m_chunks[linear_index].array; + } + + const chunk_type& chunk(point_type grid_index) const noexcept { + return chunk(m_dist.linear_index(grid_index)); + } + + chunk_type& chunk(point_type grid_index) noexcept { + return chunk(m_dist.linear_index(grid_index)); + } + + /// The home `MemoryId` of the chunk with the given row-major linear chunk index. + MemoryId chunk_home(size_t linear_index) const noexcept { + return m_chunks[linear_index].home; + } + + /// The global offset (origin) of the chunk with the given row-major linear chunk index. + point_type chunk_offset(size_t linear_index) const noexcept { + return m_dist.chunk_offset(m_dist.unravel(linear_index)); + } + + /// Returns this array restricted to `region`, as a single sub-array (in global coordinates) + /// of the one chunk that covers it. Panics if `region` does not overlap any chunk, or if it + /// spans more than one chunk -- for a region that may cross chunk boundaries, use + /// `distribution().chunk_range(region)` to enumerate the overlapping chunks yourself. + slice_type slice(bounds_type region) const { + auto grid_range = m_dist.chunk_range(region); + + if (grid_range.is_empty()) { + KMM_PANIC("`slice` region does not overlap any chunk"); + } + + for (size_t axis = 0; axis < N; axis++) { + if (grid_range.end(axis) - grid_range.begin(axis) != 1) { + KMM_PANIC("`slice` region spans multiple chunks"); + } + } + + auto linear_index = m_dist.linear_index(grid_range.begin()); + return m_chunks[linear_index].array.restrict_bounds(region); + } + + private: + struct ChunkEntry { + chunk_type array; + MemoryId home; + }; + + Distribution m_dist; + std::vector m_chunks; +}; + +} // namespace kmm diff --git a/include/kmm/api/host.hpp b/include/kmm/api/host.hpp new file mode 100644 index 00000000..4b7bbab3 --- /dev/null +++ b/include/kmm/api/host.hpp @@ -0,0 +1,78 @@ +#pragma once + +#include +#include +#include + +#include "kmm/api/context.hpp" +#include "kmm/api/resource_guard.hpp" +#include "kmm/runtime/identifiers.hpp" + +namespace kmm { + +template +class HostGuard; + +class Host: public Context { + public: + Host(Runtime runtime, MemoryTransaction transaction = {}) : + Context(std::move(runtime), std::move(transaction)) {} + + /// Acquires `args` and calls `fun(views...)`, where `views...` are the resolved views for + /// `args...`. + template + decltype(auto) submit(F&& fun, Args&&... args) { + return access(std::forward(args)...).submit(std::forward(fun)); + } + + template + HostGuard access(Args&&... args) { + return HostGuard(runtime(), transaction(), std::forward(args)...); + } +}; + +template +class HostGuard { + public: + HostGuard(Runtime& runtime, MemoryTransaction parent, Args... args) : + m_guard(runtime, MemoryId::host(), std::nullopt, parent, args...) { + runtime.synchronize(m_guard.dependencies()); + } + + ~HostGuard() { + m_guard.release(); + } + + Host context() const { + return Host(m_guard.transaction()); + } + + template + decltype(auto) get() { + return m_guard.template get(); + } + + decltype(auto) get() { + static_assert(sizeof...(Args) == 1, "argument index not specified"); + return m_guard.template get<0>(); + } + + void poison(std::exception_ptr reason) noexcept { + m_guard.poison(std::move(reason)); + } + + template + void submit(F fun) { + m_guard.apply(std::move(fun)); + } + + template + void operator>>(F fun) { + submit(std::move(fun)); + } + + private: + ResourceGuard m_guard; +}; + +} // namespace kmm diff --git a/include/kmm/api/kernel.hpp b/include/kmm/api/kernel.hpp new file mode 100644 index 00000000..91d54fd2 --- /dev/null +++ b/include/kmm/api/kernel.hpp @@ -0,0 +1,39 @@ +#pragma once + +#include + +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +/// Launcher (for use with `Device::scope`/`Device::launch`) that launches a +/// `__global__` kernel function `fun` with a fixed grid/block/shared-memory configuration. +/// `fun` is launched on `context`'s stream as +/// `fun<<>>(args...)`, where `args...` are the views +/// resolved by `scope`. +template +class Kernel { + public: + Kernel(F fun, dim3 grid_dim, dim3 block_dim, unsigned int shared_mem = 0) : + m_fun(fun), + m_grid_dim(grid_dim), + m_block_dim(block_dim), + m_shared_mem(shared_mem) {} + + template + void operator()(g_stream_t context, Args&&... args) const { + m_fun<<>>(std::forward(args)...); + } + + private: + F m_fun; + dim3 m_grid_dim; + dim3 m_block_dim; + unsigned int m_shared_mem; +}; + +/// Deduction guide: infers `F` from the kernel function passed to the constructor. +template +Kernel(F, dim3, dim3, unsigned int = 0) -> Kernel; + +} // namespace kmm diff --git a/include/kmm/api/launch_arg.hpp b/include/kmm/api/launch_arg.hpp new file mode 100644 index 00000000..312a293a --- /dev/null +++ b/include/kmm/api/launch_arg.hpp @@ -0,0 +1,137 @@ +#pragma once + +#include +#include +#include + +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/memory_manager.hpp" +#include "kmm/runtime/resource.hpp" + +namespace kmm { + +/// Tags `value` as needing read-write access when passed to `Device::scope`/`Host::scope` (the +/// default for a plain argument is read-only access). +template +struct Write { + T& value; +}; + +template +Write write(T& value) { + return Write {value}; +} + +/// Tags `value` as needing read-only access when passed to `Device::scope`/`Host::scope`. A plain +/// argument is already read-only by default; this is only useful to force a `Write`-tagged value +/// back down to read-only. +template +struct Read { + T& value; +}; + +template +Read read(T& value) { + return Read {value}; +} + +template +Read read(Write value) { + return Read {value.value}; +} + +class Runtime; + +/// Default `LaunchArg` implementation, used for any `T` without a more specific specialization +/// (e.g. a plain scalar or POD struct). The caller's value is copied once into this object; there +/// is no backing buffer to acquire or release. `LaunchArg` and `LaunchArg` build on +/// top of this copy to expose it to the kernel by const or mutable reference instead of by value. +template +class LaunchArg { + public: + using resolve_type = T; + + explicit LaunchArg(T value) : m_value(std::move(value)) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) {} + + resolve_type resolve(Runtime& runtime, const ResourceGrant& grant) { + return m_value; + } + + /// Called once after the launch's resource grant has been released (see + /// `ResourceGuard::release`). At this point the buffers touched by the launch are free again, + /// so an override may submit follow-up work on them; the default is to do nothing. + void release(Runtime& runtime) {} + + protected: + T m_value; +}; + +/// `LaunchArg` specialization used when `Device::scope`/`Host::scope` is given a const lvalue +/// reference to `T`. Just forwards to `LaunchArg`. +template +class LaunchArg: public LaunchArg { + public: + explicit LaunchArg(const T& value) : LaunchArg(value) {} +}; + +/// See `LaunchArg`. +template +class LaunchArg: public LaunchArg { + public: + explicit LaunchArg(T& value) : LaunchArg(value) {} +}; + +/// `LaunchArg` specialization for `std::tuple`. Each element `Ts` is wrapped in its own +/// `LaunchArg` (so, for example, a tuple of `Write>` elements acquires/releases each +/// array independently), and `resolve()` returns a `std::tuple` of the per-element resolved +/// values, in the same order. +template +class LaunchArg> { + public: + using resolve_type = std::tuple::resolve_type...>; + + explicit LaunchArg(std::tuple value) : m_args(std::move(value)) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) { + acquire_impl(std::index_sequence_for {}, runtime, requests, memory_id); + } + + resolve_type resolve(Runtime& runtime, const ResourceGrant& grant) { + return resolve_impl(std::index_sequence_for {}, runtime, grant); + } + + void release(Runtime& runtime) { + release_impl(std::index_sequence_for {}, runtime); + } + + private: + template + void acquire_impl( + std::index_sequence, + Runtime& runtime, + ResourceRequest& requests, + MemoryId memory_id + ) { + (std::get(m_args).acquire(runtime, requests, memory_id), ...); + } + + template + resolve_type resolve_impl( + std::index_sequence, + Runtime& runtime, + const ResourceGrant& grant + ) { + return resolve_type {std::get(m_args).resolve(runtime, grant)...}; + } + + template + void release_impl(std::index_sequence, Runtime& runtime) { + (std::get(m_args).release(runtime), ...); + } + + std::tuple...> m_args; +}; + +} // namespace kmm diff --git a/include/kmm/api/launcher.hpp b/include/kmm/api/launcher.hpp deleted file mode 100644 index 528537d1..00000000 --- a/include/kmm/api/launcher.hpp +++ /dev/null @@ -1,80 +0,0 @@ -#pragma once - -#include "kmm/core/domain.hpp" -#include "kmm/core/resource.hpp" - -namespace kmm { - -template -struct Host { - static constexpr ExecutionSpace execution_space = ExecutionSpace::Host; - - Host(F fun) : m_fun(fun) {} - - template - void operator()(Resource& resource, DomainChunk chunk, Args... args) { - m_fun(args...); - } - - private: - std::decay_t m_fun; -}; - -template -struct GPU { - static constexpr ExecutionSpace execution_space = ExecutionSpace::Device; - - GPU(F fun) : m_fun(fun) {} - - template - void operator()(Resource& resource, DomainChunk chunk, Args... args) { - m_fun(resource.cast(), args...); - } - - private: - std::decay_t m_fun; -}; - -template -struct GPUKernel { - static constexpr ExecutionSpace execution_space = ExecutionSpace::Device; - - GPUKernel(F kernel, dim3 block_size) : GPUKernel(kernel, block_size, block_size) {} - - GPUKernel(F kernel, dim3 block_size, dim3 elements_per_block, uint32_t shared_memory = 0) : - kernel(kernel), - block_size(block_size), - elements_per_block(elements_per_block), - shared_memory(shared_memory) {} - - template - void operator()(Resource& resource, DomainChunk chunk, Args... args) { - int64_t g[3] = { - chunk.size.get_or_default(0), - chunk.size.get_or_default(1), - chunk.size.get_or_default(2) - }; - int64_t b[3] = {elements_per_block.x, elements_per_block.y, elements_per_block.z}; - dim3 grid_dim = { - checked_cast((g[0] / b[0]) + int64_t(g[0] % b[0] != 0)), - checked_cast((g[1] / b[1]) + int64_t(g[1] % b[1] != 0)), - checked_cast((g[2] / b[2]) + int64_t(g[2] % b[2] != 0)), - }; - - resource.cast().launch( // - grid_dim, - block_size, - shared_memory, - kernel, - args... - ); - } - - private: - std::decay_t kernel; - dim3 block_size; - dim3 elements_per_block; - uint32_t shared_memory; -}; - -} // namespace kmm diff --git a/include/kmm/api/mapper.hpp b/include/kmm/api/mapper.hpp deleted file mode 100644 index d1eeb7d8..00000000 --- a/include/kmm/api/mapper.hpp +++ /dev/null @@ -1,268 +0,0 @@ -#pragma once - -#include "spdlog/spdlog.h" - -#include "kmm/api/argument.hpp" -#include "kmm/core/domain.hpp" -#include "kmm/core/reduction.hpp" -#include "kmm/utils/geometry.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -struct All { - template - Bounds operator()(DomainChunk chunk, Bounds bounds) const { - return bounds; - } -}; - -struct Axis { - constexpr Axis() : m_axis(0) {} - explicit constexpr Axis(size_t axis) : m_axis(axis) {} - - Bounds<1> operator()(DomainChunk chunk) const { - return Bounds<1>::from_offset_size( - chunk.offset.get_or_default(m_axis), - chunk.size.get_or_default(m_axis) - ); - } - - Bounds<1> operator()(DomainChunk chunk, Bounds<1> bounds) const { - return (*this)(chunk).intersection(bounds); - } - - size_t get() const { - return m_axis; - } - - explicit operator size_t() const { - return get(); - } - - private: - size_t m_axis = 0; -}; - -struct IdentityMap { - template - Bounds operator()(DomainChunk chunk, Bounds bounds) const { - return Bounds::from_offset_size(Point::from(chunk.offset), Dim::from(chunk.size)); - } -}; - -// (scale * variable + offset + [0...length]) / divisor -struct IndexMap { - constexpr IndexMap(Axis variable = {}) : m_axis(variable) {} - - IndexMap( - Axis variable, - int64_t scale, - int64_t offset = 0, - int64_t length = 1, - int64_t divisor = 1 - ); - - static IndexMap range(IndexMap begin, IndexMap end); - IndexMap offset_by(int64_t offset) const; - IndexMap scale_by(int64_t factor) const; - IndexMap divide_by(int64_t divisor) const; - IndexMap negate() const; - Bounds<1> apply(DomainChunk chunk) const; - - Bounds<1> operator()(DomainChunk chunk) const { - return apply(chunk); - } - - Bounds<1> operator()(DomainChunk chunk, Bounds<1> bounds) const { - return apply(chunk).intersection(bounds); - } - - friend std::ostream& operator<<(std::ostream& f, const IndexMap& that); - - private: - Axis m_axis = {}; - int64_t m_scale = 1; - int64_t m_offset = 0; - int64_t m_length = 1; - int64_t m_divisor = 1; -}; - -inline IndexMap range(IndexMap begin, IndexMap end) { - return IndexMap::range(begin, end); -} - -inline IndexMap range(int64_t begin, int64_t end) { - return {Axis {}, 0, begin, end - begin}; -} - -inline IndexMap range(int64_t end) { - return range(0, end); -} - -inline IndexMap operator+(IndexMap a) { - return a; -} - -inline IndexMap operator+(IndexMap a, int64_t b) { - return a.offset_by(b); -} - -inline IndexMap operator+(int64_t a, IndexMap b) { - return b.offset_by(a); -} - -inline IndexMap operator-(IndexMap a) { - return a.negate(); -} - -inline IndexMap operator-(IndexMap a, int64_t b) { - return a + (-b); -} - -inline IndexMap operator-(int64_t a, IndexMap b) { - return a + (-b); -} - -inline IndexMap operator*(IndexMap a, int64_t b) { - return a.scale_by(b); -} - -inline IndexMap operator*(int64_t a, IndexMap b) { - return b.scale_by(a); -} - -inline IndexMap operator/(IndexMap a, int64_t b) { - return a.divide_by(b); -} - -template -struct MultiIndexMap { - Bounds operator()(DomainChunk chunk) const { - Bounds result; - - for (size_t i = 0; i < N; i++) { - result[i] = (this->axes[i])(chunk); - } - - return result; - } - - Bounds operator()(DomainChunk chunk, Bounds bounds) const { - Bounds result; - - for (size_t i = 0; i < N; i++) { - result[i] = (this->axes[i])(chunk, Bounds<1> {bounds[i]}); - } - - return result; - } - - IndexMap axes[N]; -}; - -template<> -struct MultiIndexMap<0> { - Bounds<0> operator()(DomainChunk chunk, Bounds<0> bounds = {}) const { - return {}; - } -}; - -inline IndexMap into_index_map(int64_t m) { - return {Axis {}, 0, m}; -} - -inline IndexMap into_index_map(Axis m) { - return m; -} - -inline IndexMap into_index_map(IndexMap m) { - return m; -} - -inline IndexMap into_index_map(All m) { - return {Axis(), 0, 0, std::numeric_limits::max()}; -} - -template -MultiIndexMap bounds(const Is&... slices) { - return {into_index_map(slices)...}; -} - -template -MultiIndexMap tile(const Is&... length) { - size_t variable = 0; - return {IndexMap( - Axis {variable++}, - checked_cast(length), - 0, - checked_cast(length) - )...}; -} - -namespace placeholders { -static constexpr All _; - -static constexpr Axis _x = Axis(0); -static constexpr Axis _y = Axis(1); -static constexpr Axis _z = Axis(2); - -static constexpr Axis _i = Axis(0); -static constexpr Axis _j = Axis(1); -static constexpr Axis _k = Axis(2); - -static constexpr Axis _0 = Axis(0); -static constexpr Axis _1 = Axis(1); -static constexpr Axis _2 = Axis(2); - -static constexpr MultiIndexMap<2> _xy = {_x, _y}; -static constexpr MultiIndexMap<3> _xyz = {_x, _y, _z}; - -static constexpr MultiIndexMap<2> _ij = {_x, _y}; -static constexpr MultiIndexMap<3> _ijk = {_x, _y, _z}; - -static constexpr IdentityMap one_to_one; -static constexpr All all; -} // namespace placeholders - -template<> -struct Argument: Argument> { - static Argument pack(TaskInstance& builder, IndexMap mapper) { - return {mapper(builder.chunk).get_or_default(0)}; - } -}; - -template<> -struct Argument: Argument> { - static Argument pack(TaskInstance& builder, Axis mapper) { - return {mapper(builder.chunk).get_or_default(0)}; - } -}; - -template -struct Argument>: Argument> { - static Argument pack(TaskInstance& builder, MultiIndexMap mapper) { - return {mapper(builder.chunk)}; - } -}; - -namespace detail { -template -struct RangeDim: std::integral_constant {}; - -template -struct RangeDim>: std::integral_constant {}; -} // namespace detail - -template -static constexpr size_t mapper_dimensionality = - detail::RangeDim>::value; - -template -static constexpr bool is_dimensionality_accepted_by_mapper = - detail::RangeDim>>::value == N; - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/api/parallel_for.hpp b/include/kmm/api/parallel_for.hpp new file mode 100644 index 00000000..c4859d03 --- /dev/null +++ b/include/kmm/api/parallel_for.hpp @@ -0,0 +1,110 @@ +#pragma once + +#include +#include + +#include "kmm/api/kernel.hpp" +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/shape.hpp" + +namespace kmm { + +namespace detail { +template +__global__ void parallel_for_kernel( // + IndexMapper mapper, F fun, Args... args) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Point delta; + + if (mapper.unravel(linear_index, delta)) { + fun(Point::from(delta), args...); + } +} + +template +__global__ void parallel_for_offset_kernel( + Point offset, + IndexMapper mapper, + F fun, + Args... args +) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Point delta; + + if (mapper.unravel(linear_index, delta)) { + fun(offset + Point::from(delta), args...); + } +} +} // namespace detail + +/// Launcher (for use with `Device::access`/`DeviceGuard::parallel_for`) that applies `fun` to +/// every point of the N-dimensional index space `shape`. +template +class ParallelFor { + public: + using index_type = default_index_type; + + explicit ParallelFor(Shape shape, F fun, unsigned int block_size = 256) : + m_shape(shape), + m_fun(std::move(fun)), + m_block_size(block_size) {} + + template + void operator()(g_stream_t stream, Args&&... args) const { + launch_recur(stream, m_offset, m_shape, args...); + } + + private: + template + void launch_recur( + g_stream_t stream, + Point offset, + Shape shape, + const Args&... args + ) const { + if (shape.is_empty()) { + return; + } + + uint32_t num_threads = 1; + + for (size_t i = 0; i < N; i++) { + // chunk is the maximum number of elements that can be along the i-th + // dimensions without `num_threads` exceeding its maximum. + auto chunk = index_type(IndexMapper::max_volume / num_threads); + + // if chunk < shape[i], then we must split the i-th dimensions. + // we strip off the first `chunk` elements and launch it recursively. + while (is_less(chunk, shape[i])) { + Shape head = shape; + head[i] = chunk; + launch_recur(stream, offset, head, args...); + + offset[i] += chunk; + shape[i] -= chunk; + } + + num_threads *= uint32_t(shape[i]); + } + + uint32_t grid_size = div_ceil(num_threads, m_block_size); + auto mapper = IndexMapper(Shape::from(shape)); + + if (offset == Point::zero()) { + detail::parallel_for_kernel...> + <<>>(mapper, m_fun, args...); + } else { + detail::parallel_for_offset_kernel...> + <<>>(offset, mapper, m_fun, args...); + } + } + + F m_fun; + Point m_offset; + Shape m_shape; + uint32_t m_block_size; +}; + +} // namespace kmm diff --git a/include/kmm/api/parallel_reduce.hpp b/include/kmm/api/parallel_reduce.hpp new file mode 100644 index 00000000..26611f9c --- /dev/null +++ b/include/kmm/api/parallel_reduce.hpp @@ -0,0 +1,177 @@ +#pragma once + +#include +#include + +#ifdef KMM_USE_CUDA + #include "cub/block/block_reduce.cuh" +#elif KMM_USE_HIP + #include "hipcub/hipcub.hpp" +namespace cub = hipcub; +#endif + +#include "kmm/api/device.hpp" +#include "kmm/api/kernel.hpp" +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/shape.hpp" + +namespace kmm { + +namespace detail { +template +__global__ void parallel_reduce_kernel( // + IndexMapper mapper, F fun, OutputT* output_addr, Args... args) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Point delta; + OutputT output {}; + + if (mapper.unravel(linear_index, delta)) { + output += fun(Point::from(delta), args...); + } + + output = cub::BlockReduce().Sum(output); + + if (threadIdx.x == 0) { + output_addr[blockIdx.x] = output; + } +} + +template +__global__ void parallel_reduce_offset_kernel( + Point offset, + IndexMapper mapper, + F fun, + OutputT* output_addr, + Args... args +) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Point delta; + OutputT output {}; + + if (mapper.unravel(linear_index, delta)) { + output += fun(offset + Point::from(delta), args...); + } + + output = cub::BlockReduce().Sum(output); + + if (threadIdx.x == 0) { + output_addr[blockIdx.x] = output; + } +} +} // namespace detail + +/// Launcher (for use with `Device::access`/`DeviceGuard::parallel_for`) that applies `fun` to +/// every point of the N-dimensional index space `shape`. +template +class ParallelReduce { + public: + using index_type = default_index_type; + + explicit ParallelReduce(Shape shape, F fun, unsigned int block_size = 256) : + m_shape(shape), + m_fun(std::move(fun)), + m_block_size(block_size) {} + + template + void operator()(g_stream_t stream, kmm::ViewMut output, Args&&... args) const { + launch_recur(stream, m_offset, m_shape, output.data(), args...); + } + + template + void operator()(g_stream_t stream, OutputT* output_addr, Args&&... args) const { + launch_recur(stream, m_offset, m_shape, output_addr, args...); + } + + size_t num_outputs() const { + return num_outputs_impl(m_shape); + } + + private: + size_t num_outputs_impl(Shape shape) const { + if (shape.is_empty()) { + return 0; + } + + size_t num_outputs = 0; + uint32_t num_threads = 1; + + for (size_t i = 0; i < N; i++) { + auto chunk = index_type(IndexMapper::max_volume / num_threads); + + while (is_less(chunk, shape[i])) { + Shape head = shape; + head[i] = chunk; + num_outputs += num_outputs_impl(head); + + shape[i] -= chunk; + } + + num_threads *= uint32_t(shape[i]); + } + + uint32_t grid_size = div_ceil(num_threads, m_block_size); + return num_outputs + grid_size; + } + + template + OutputT* launch_recur( + g_stream_t stream, + Point offset, + Shape shape, + OutputT* output_addr, + const Args&... args + ) const { + if (shape.is_empty()) { + return output_addr; + } + + uint32_t num_threads = 1; + + for (size_t i = 0; i < N; i++) { + // chunk is the maximum number of elements that can be along the i-th + // dimensions without `num_threads` exceeding its maximum. + auto chunk = index_type(IndexMapper::max_volume / num_threads); + + // if chunk < shape[i], then we must split the i-th dimensions. + // we strip off the first `chunk` elements and launch it recursively. + while (is_less(chunk, shape[i])) { + Shape head = shape; + head[i] = chunk; + output_addr = launch_recur(stream, offset, head, output_addr, args...); + + offset[i] += chunk; + shape[i] -= chunk; + } + + num_threads *= uint32_t(shape[i]); + } + + uint32_t grid_size = div_ceil(num_threads, m_block_size); + auto mapper = IndexMapper(Shape::from(shape)); + + if (offset == Point::zero()) { + detail::parallel_reduce_kernel...> + <<>>(mapper, m_fun, output_addr, args...); + } else { + detail::parallel_reduce_offset_kernel...> + <<>>( + offset, + mapper, + m_fun, + output_addr, + args... + ); + } + + return output_addr + grid_size; + } + + F m_fun; + Point m_offset; + Shape m_shape; + uint32_t m_block_size; +}; + +} // namespace kmm diff --git a/include/kmm/api/parallel_submit.hpp b/include/kmm/api/parallel_submit.hpp deleted file mode 100644 index cd5645eb..00000000 --- a/include/kmm/api/parallel_submit.hpp +++ /dev/null @@ -1,130 +0,0 @@ -#pragma once - -#include "kmm/api/argument.hpp" -#include "kmm/api/task_group.hpp" -#include "kmm/core/buffer.hpp" -#include "kmm/core/domain.hpp" -#include "kmm/core/identifiers.hpp" -#include "kmm/runtime/runtime.hpp" - -namespace kmm { - -class Runtime; -class TaskGraphState; - -namespace detail { - -template -class ComputeTaskImpl: public ComputeTask { - public: - ComputeTaskImpl(DomainChunk chunk, Launcher launcher, Args... args) : - m_chunk(chunk), - m_launcher(std::move(launcher)), - m_args(std::move(args)...) {} - - void execute(Resource& resource, TaskContext context) override { - execute_impl(std::index_sequence_for(), resource, context); - } - - template - void execute_impl(std::index_sequence, Resource& resource, TaskContext& context) { - static constexpr ExecutionSpace execution_space = Launcher::execution_space; - - m_launcher( - resource, - m_chunk, - ArgumentUnpack::call(context, std::get(m_args))... - ); - } - - private: - DomainChunk m_chunk; - Launcher m_launcher; - std::tuple m_args; -}; - -template -EventId parallel_submit_impl( - std::index_sequence, - Runtime& runtime, - const SystemInfo& system_info, - const Domain& domain, - Launcher launcher, - Args&&... args -) { - std::tuple...> handlers = {std::forward(args)...}; - - auto init = TaskGroupInit { - .runtime = runtime, // - .domain = domain - }; - - (std::get(handlers).initialize(init), ...); - - return runtime.schedule([&](TaskGraph& graph) { - EventList events; - - for (const DomainChunk& chunk : domain.chunks) { - auto processor_id = chunk.owner_id; - - auto instance = TaskInstance { - .runtime = runtime, - .graph = graph, - .chunk = chunk, - .memory_id = system_info.affinity_memory(processor_id), - .buffers = {}, - .dependencies = {} - }; - - auto task = std::make_unique...>>( - chunk, - launcher, - std::get(handlers).before_submit(instance)... - ); - - EventId event_id = graph.insert_compute_task( - processor_id, - std::move(task), - std::move(instance.buffers), - std::move(instance.dependencies) - ); - - events.push_back(event_id); - - auto result = TaskSubmissionResult { - .runtime = runtime, // - .graph = graph, - .event_id = event_id - }; - - (std::get(handlers).after_submit(result), ...); - } - - auto commit = TaskGroupCommit {.runtime = runtime, .graph = graph}; - - (std::get(handlers).commit(commit), ...); - - return graph.join_events(events); - }); -} -} // namespace detail - -template -EventId parallel_submit( - Runtime& runtime, - const SystemInfo& system_info, - const Domain& partition, - Launcher launcher, - Args&&... args -) { - return detail::parallel_submit_impl( - std::index_sequence_for {}, - runtime, - system_info, - partition, - launcher, - std::forward(args)... - ); -} - -} // namespace kmm diff --git a/include/kmm/api/reduce.hpp b/include/kmm/api/reduce.hpp new file mode 100644 index 00000000..865bddd0 --- /dev/null +++ b/include/kmm/api/reduce.hpp @@ -0,0 +1,246 @@ +#pragma once + +#include +#include +#include + +#include "kmm/api/accumulator.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/runtime.hpp" + +namespace kmm { + +template +struct Reduce { + const T& target; + ReductionOp op; + Shape replication; +}; + +template +Reduce, sizeof...(Extents)> reduce( + const NDArray& array, + ReductionOp op, + Extents... extents +) { + auto replication = Shape {checked_cast(extents)...}; + + return Reduce, sizeof...(Extents)> {array, op, replication}; +} + +template +Reduce, sizeof...(Extents)> reduce( + const NDArray& array, + Extents... extents +) { + return reduce(array, ReductionOp::Sum, extents...); +} + +template +Reduce, sizeof...(Extents)> reduce( + const NDAccumulator& accumulator, + Extents... extents +) { + auto replication = Shape {checked_cast(extents)...}; + + return Reduce, sizeof...(Extents)> { + accumulator, + accumulator.op(), + replication + }; +} + +namespace detail { + +template +class ReplicatedRegion { + public: + using index_type = typename LayoutT::index_type; + static constexpr size_t rank = LayoutT::rank; + static constexpr size_t replicated_rank = rank + K; + using replicated_shape = Shape; + using replicated_strides = StridesN; + using replicated_layout = Layout; + + ReplicatedRegion(const LayoutT& layout, const Shape& replication) : + m_source(layout.normalize_offset()), + m_replication(replication) { + KMM_ASSERT(m_replication.volume() >= 1); + + auto span = m_source.offset_span(); + m_replica_span = checked_sub(span.stop, span.start); + } + + size_t replica_span() const { + return m_replica_span; + } + + size_t replica_count() const { + return checked_cast(m_replication.volume()); + } + + size_t element_count() const { + return checked_mul(replica_span(), replica_count()); + } + + replicated_layout view() const { + auto array_shape = m_source.shape(); + auto array_strides = m_source.strides(); + + replicated_shape shape; + Vec strides; + + if constexpr (rank > 0) { + for (size_t i = 0; i < rank; i++) { + shape[i] = array_shape[i]; + strides[i] = static_cast(array_strides[i]); + } + } + + auto stride = static_cast(replica_span()); + for (size_t j = K; j-- > 0;) { + shape[rank + j] = m_replication[j]; + strides[rank + j] = stride; + stride *= static_cast(m_replication[j]); + } + + auto mapping = build_strides(strides, make_index_sequence()); + return replicated_layout(shape, mapping, m_source.base_offset()); + } + + ReductionDescription fold(ReductionOp op, DataType dtype) const { + auto elem_size = checked_cast(data_type_size(dtype)); + auto base_offset = checked_mul(m_source.base_offset(), elem_size); + + ReductionDescription description(dtype, op); + description.input_offset = base_offset; + description.output_offset = base_offset; + + if constexpr (rank > 0) { + auto array_shape = m_source.shape(); + auto array_strides = m_source.strides(); + + for (size_t i = 0; i < rank; i++) { + auto stride_bytes = checked_mul(array_strides[i], elem_size); + description.add_dimension( + checked_cast(array_shape[i]), + stride_bytes, + stride_bytes + ); + } + } + + description.reduction_extent = checked_cast(replica_count()); + description.reduction_stride = checked_mul(replica_span(), elem_size); + return description; + } + + private: + template + KMM_HOST_DEVICE static replicated_strides build_strides( + const Vec& values, + IndexSequence + ) { + return replicated_strides {values[Is]...}; + } + + LayoutT m_source; + Shape m_replication; + size_t m_replica_span; +}; + +template +class LaunchArgReduce { + public: + using index_type = typename LayoutT::index_type; + using region_type = ReplicatedRegion; + using replicated_layout = typename region_type::replicated_layout; + using resolve_type = NDView; + + LaunchArgReduce( + Buffer buffer, + LayoutT layout, + ReductionOp op, + Shape replication + ) : + m_buffer(std::move(buffer)), + m_region(layout, replication), + m_op(op) {} + + void acquire(Runtime& runtime, ResourceRequest& requests, MemoryId memory_id) { + m_memory_id = memory_id; + + m_scratch = runtime.create_buffer( + BufferLayout::for_type(m_region.element_count()), + "replicated partial", + reduction_identity(data_type_of(), m_op) + ); + + m_index = requests.add(memory_id, *m_scratch, AccessMode::ReadWrite); + } + + resolve_type resolve(Runtime& runtime, const ResourceGrant& grant) { + auto accessor = grant.accessor(m_index); + auto layout = m_region.view(); + return {static_cast(accessor.address), layout}; + } + + void release(Runtime& runtime) { + if (!m_scratch) { + return; + } + + runtime.submit_reduction( + m_buffer.id(), + *m_scratch, + m_region.fold(m_op, data_type_of()), + m_memory_id + ); + runtime.release_buffer(*m_scratch); + m_scratch.reset(); + } + + private: + Buffer m_buffer; + region_type m_region; + ReductionOp m_op; + MemoryId m_memory_id = MemoryId::host(); + std::optional m_scratch; + size_t m_index = 0; +}; + +} // namespace detail + +template +class LaunchArg, K>>: public detail::LaunchArgReduce { + public: + explicit LaunchArg(const Reduce, K>& arg) : + detail::LaunchArgReduce( + arg.target.buffer(), + arg.target.layout(), + arg.op, + arg.replication + ) {} +}; + +template +class LaunchArg, K>>: + public detail::LaunchArgReduce { + public: + explicit LaunchArg(const Reduce, K>& arg) : + detail::LaunchArgReduce( + arg.target.buffer(), + arg.target.layout(), + arg.op, + arg.replication + ) {} +}; + +template +class LaunchArg, 0>>: public LaunchArg> { + public: + explicit LaunchArg(const Reduce, 0>& arg) : + LaunchArg>(arg.target) {} +}; + +} // namespace kmm diff --git a/include/kmm/api/resource_guard.hpp b/include/kmm/api/resource_guard.hpp new file mode 100644 index 00000000..c71928cb --- /dev/null +++ b/include/kmm/api/resource_guard.hpp @@ -0,0 +1,147 @@ +#pragma once + +#include +#include +#include +#include + +#include "kmm/api/launch_arg.hpp" +#include "kmm/runtime/resource.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +template +class ResourceGuard { + public: + ResourceGuard( + Runtime& runtime, + MemoryId memory_id, + std::optional stream, + MemoryTransaction parent, + Args... args + ) { + std::tuple...> args_tuple(args...); + + // `ResourceGrant` is not movable, so it must be initialized in place. + m_state = std::unique_ptr(new State { + runtime, + submit_impl( + std::index_sequence_for {}, + args_tuple, + runtime, + memory_id, + stream, + std::move(parent) + ), + std::move(args_tuple) + }); + } + + Runtime& runtime() const noexcept { + return m_state->runtime; + } + + const MemoryTransaction& transaction() const noexcept { + return m_state->grant.transaction(); + } + + const DeviceEventSet& dependencies() const noexcept { + return m_state->grant.dependencies(); + } + + template + decltype(auto) get() { + return std::get(m_state->args).resolve(m_state->runtime, m_state->grant); + } + + template + decltype(auto) apply(F&& fun, Extra&&... extra) { + try { + return apply_impl( + std::index_sequence_for {}, + std::forward(fun), + std::forward(extra)... + ); + } catch (...) { + poison(std::current_exception()); + throw; + } + } + + void release(DeviceEventSet deps = {}) { + std::exception_ptr error; + + // The grant is released first so that each arg's `release()` runs as a genuine + // post-completion hook: by the time it is called the buffers this launch touched are no + // longer held by the grant, so `release()` may itself submit follow-up work on them + // (e.g. `reduce(array, k)` folding its scratch buffer into the destination array). + try { + m_state->runtime.release(m_state->grant, std::move(deps)); + } catch (...) { + error = std::current_exception(); + } + + // Every arg's `release()` must still run even when an earlier step threw, otherwise a + // single failure would strand the buffers held by the remaining args. The first error + // seen (from the grant release or any arg) is rethrown once every arg has had its turn. + release_impl(std::index_sequence_for {}, error); + + if (error) { + std::rethrow_exception(error); + } + } + + void poison(std::exception_ptr reason) noexcept { + m_state->runtime.poison(m_state->grant, reason); + } + + private: + struct State { + Runtime runtime; + ResourceGrant grant; + std::tuple...> args; + }; + + template + static ResourceGrant submit_impl( + std::index_sequence, + std::tuple...>& args, + Runtime& runtime, + MemoryId memory_id, + std::optional stream, + MemoryTransaction parent + ) { + ResourceRequest requests; + (std::get(args).acquire(runtime, requests, memory_id), ...); + return runtime.submit(std::move(requests), stream, std::move(parent)); + } + + template + decltype(auto) apply_impl(std::index_sequence, F&& fun, Extra&&... extra) { + return fun( + std::forward(extra)..., + std::get(m_state->args).resolve(m_state->runtime, m_state->grant)... + ); + } + + template + void release_impl(std::index_sequence, std::exception_ptr& error) { + (release_one(std::get(m_state->args), error), ...); + } + + template + void release_one(Arg& arg, std::exception_ptr& error) { + try { + arg.release(m_state->runtime); + } catch (...) { + if (!error) { + error = std::current_exception(); + } + } + } + + std::unique_ptr m_state; +}; + +} // namespace kmm diff --git a/include/kmm/api/runtime_handle.hpp b/include/kmm/api/runtime_handle.hpp deleted file mode 100644 index 78346200..00000000 --- a/include/kmm/api/runtime_handle.hpp +++ /dev/null @@ -1,266 +0,0 @@ -#pragma once - -#include -#include - -#include "kmm/api/array.hpp" -#include "kmm/api/launcher.hpp" -#include "kmm/api/parallel_submit.hpp" -#include "kmm/api/struct_argument.hpp" -#include "kmm/core/config.hpp" -#include "kmm/core/system_info.hpp" -#include "kmm/core/view.hpp" -#include "kmm/utils/checked_math.hpp" -#include "kmm/utils/panic.hpp" -#include "kmm/utils/range.hpp" - -namespace kmm { - -class Runtime; - -class RuntimeHandle { - struct Impl; - RuntimeHandle(std::shared_ptr impl); - - public: - RuntimeHandle(std::shared_ptr rt); - RuntimeHandle(Runtime& rt); - - /** - * Submit a single task to the runtime system. - * - * @param index_space The index space defining the task dimensions. - * @param target The target processor for the task. - * @param launcher The task launcher. - * @param args The arguments that are forwarded to the launcher. - * @return The event identifier for the submitted task. - */ - template - EventId submit(ResourceId target, L&& launcher, Args&&... args) const { - DomainChunk chunk = { - .owner_id = target, // - .offset = DomainPoint::zero(), - .size = DomainDim::one() - }; - - return kmm::parallel_submit( - worker(), - info(), - Domain {{chunk}}, - std::forward(launcher), - std::forward(args)... - ); - } - - /** - * Submit a set of tasks to the runtime systems. - * - * @param dist The domain describing how the work is split. - * @param launcher The task launcher. - * @param args The arguments that are forwarded to the launcher. - * @return The event identifier for the submitted task. - */ - template - EventId parallel_submit(D&& domain, L&& launcher, Args&&... args) const { - return kmm::parallel_submit( - worker(), - info(), - IntoDomain>::call( - std::forward(domain), - info(), - std::decay_t::execution_space - ), - std::forward(launcher), - std::forward(args)... - ); - } - - /** - * Submit a set of tasks to the runtime systems. - * - * @param domain_size The index space defining the domain dimensions. - * @param partitioner The partitioner describing how the work is split. - * @param launcher The task launcher. - * @param args The arguments that are forwarded to the launcher. - * @return The event identifier for the submitted task. - */ - template - EventId parallel_submit( - DomainDim domain_size, - DomainDim chunk_size, - L&& launcher, - Args&&... args - ) const { - return this->parallel_submit( - TileDomain(domain_size, chunk_size), - std::forward(launcher), - std::forward(args)... - ); - } - - /** - * Allocates an array in memory with the given shape and memory affinity. - * - * The pointer to the given buffer should contain `shape[0] * shape[1] * shape[2]...` - * elements. - * - * @param data Pointer to the array data. - * @param shape Shape of the array. - * @param memory_id Identifier of the memory region. - * @return The allocated Array object. - */ - template - Array allocate(const T* data, Dim shape, MemoryId memory_id) const { - auto handle = ArrayInstance::create( - worker(), - Distribution {shape, shape, {memory_id}}, - DataType::of() - ); - - handle->copy_bytes_from(data); - return Array {std::move(handle)}; - } - - /** - * Allocates an array in memory with the given shape. - * - * The pointer to the given buffer should contain `shape[0] * shape[1] * shape[2]...` - * elements. - * - * In which memory the data will be allocated is determined by `memory_affinity_for_address`. - * - * @param data Pointer to the array data. - * @param shape Shape of the array. - * @return The allocated Array object. - */ - template - Array allocate(const T* data, Dim shape) const { - return allocate(data, shape, memory_affinity_for_address(data)); - } - - /** - * Alias for `allocate(v.data(), v.sizes())` - */ - template - Array allocate(View v) const { - return allocate(v.data(), v.sizes()); - } - - /** - * Alias for `allocate(data, Dim{sizes...})` - */ - template - Array allocate(const T* data, const Is&... num_elements) const { - return allocate(data, Dim {checked_cast(num_elements)...}); - } - - /** - * Alias for `allocate(v.data(), v.size())` - */ - template - Array allocate(const std::vector& v) const { - return allocate(v.data(), v.size()); - } - - /** - * Alias for `allocate(v.begin(), v.size())` - */ - template - Array allocate(std::initializer_list v) const { - return allocate(v.begin(), v.size()); - } - - /** - * Returns the memory affinity for a given address. - */ - MemoryId memory_affinity_for_address(const void* address) const; - - /** - * Returns a new event that triggers when all the given events have triggered. - */ - EventId join(EventList events) const; - - /** - * Returns a new event that triggers when all the given events have triggered. Each argument - * must be convertible to an `EventId`. - */ - template - EventId join(Es... events) const { - return join(EventList(EventId(events)...)); - } - - /** - * Returns `true` if the event with the provided identifier has finished, or `false` otherwise. - */ - bool is_done(EventId) const; - - /** - * Block the current thread until the event with the provided identifier completes. - */ - void wait(EventId id) const; - - /** - * Block the current thread until the event with the provided id completes. Blocks until - * either the event completes or the deadline is exceeded, whatever comes first. - * - * @return `true` if the event with the provided id has finished, otherwise returns `false`. - */ - bool wait_until(EventId id, std::chrono::system_clock::time_point deadline) const; - - /** - * Block the current thread until the event with the provided id completes. Blocks until - * either the event completes or the duration is exceeded, whatever comes first. - * - * @return `true` if the event with the provided id has finished, otherwise returns `false`. - */ - bool wait_for(EventId id, std::chrono::system_clock::duration duration) const; - - /** - * Submit a barrier the runtime system. The barrier completes once all the tasks submitted - * to the runtime system so far have finished. - * - * @return The identifier of the barrier. - */ - EventId barrier() const; - - /** - * Blocks until all the tasks submitted to the runtime system have finished and the - * system has become idle. - */ - void synchronize() const; - - /** - * Return a new `RuntimeHandle` that is constrained to the given set of resources. In other - * words, it only can submit work onto those resources. - */ - RuntimeHandle constrain_to(std::vector resources) const; - - /** - * Return a new `RuntimeHandle` that is constrained to the given device. In other - * words, it only can submit work onto that device. - */ - RuntimeHandle constrain_to(DeviceId device) const; - - /** - * Return a new `RuntimeHandle` that is constrained to the given resource. In other - * words, it only can submit work onto that resource. - */ - RuntimeHandle constrain_to(ResourceId resource) const; - - /** - * Returns information about the current system. - */ - const SystemInfo& info() const; - - /** - * Returns the inner `Worker`. - */ - Runtime& worker() const; - - private: - std::shared_ptr m_data; -}; - -RuntimeHandle make_runtime(const RuntimeConfig& config = default_config_from_environment()); - -} // namespace kmm diff --git a/include/kmm/api/struct_argument.hpp b/include/kmm/api/struct_argument.hpp deleted file mode 100644 index 5e3b2381..00000000 --- a/include/kmm/api/struct_argument.hpp +++ /dev/null @@ -1,94 +0,0 @@ -#pragma once - -#include "kmm/api/argument.hpp" - -namespace kmm { - -template -struct StructArgument { - static constexpr size_t num_fields = sizeof...(Fields); - std::tuple fields; -}; - -template -struct StructArgumentHandler { - using type = StructArgument::type...>; - - StructArgumentHandler(const Type&, Fields... fields) : m_handlers(fields...) {} - - void initialize(const TaskGroupInit& init) { - initialize_impl(init, std::make_index_sequence()); - } - - type before_submit(TaskInstance& task) { - return before_submit_impl(task, std::make_index_sequence()); - } - - void after_submit(const TaskSubmissionResult& result) { - after_submit_impl(result, std::make_index_sequence()); - } - - void commit(const TaskGroupCommit& commit) { - commit_impl(commit, std::make_index_sequence()); - } - - private: - template - void initialize_impl(const TaskGroupInit& init, std::index_sequence) { - (std::get(m_handlers).initialize(init), ...); - } - - template - type before_submit_impl(TaskInstance& task, std::index_sequence) { - return {.fields = {(std::get(m_handlers).before_submit(task))...}}; - } - - template - void after_submit_impl(const TaskSubmissionResult& result, std::index_sequence) { - (std::get(m_handlers).after_submit(result), ...); - } - - template - void commit_impl(const TaskGroupCommit& commit, std::index_sequence) { - (std::get(m_handlers).commit(commit), ...); - } - - std::tuple...> m_handlers; -}; - -template -struct StructArgumentUnpack { - static View call(TaskContext& context, StructArgument& data) { - return call_impl(context, data, std::index_sequence_for()); - } - - private: - template - static View - call_impl(TaskContext& context, StructArgument& data, std::index_sequence) { - return {ArgumentUnpack::call(context, std::get(data.fields))...}; - } -}; - -} // namespace kmm - -#define KMM_DEFINE_STRUCT_ARGUMENT_IMPL(UNIQUE_NAME, T, ...) \ - static auto UNIQUE_NAME(const T& it) { \ - return kmm::StructArgumentHandler(it, __VA_ARGS__); \ - } \ - template<> \ - struct kmm::ArgumentHandler: decltype(UNIQUE_NAME(std::declval())) { \ - ArgumentHandler(const T& it) : decltype(UNIQUE_NAME(it))(it, __VA_ARGS__) {} \ - }; - -#define KMM_DEFINE_STRUCT_ARGUMENT(T, ...) \ - KMM_DEFINE_STRUCT_ARGUMENT_IMPL( \ - KMM_CONCAT(__kmm_argument_handle_type_helper_, __LINE__), \ - T, \ - __VA_ARGS__ \ - ) - -#define KMM_DEFINE_STRUCT_VIEW(T, V) \ - template \ - struct kmm::ArgumentUnpack>: \ - kmm::StructArgumentUnpack {}; diff --git a/include/kmm/api/task_group.hpp b/include/kmm/api/task_group.hpp deleted file mode 100644 index 5bf7492c..00000000 --- a/include/kmm/api/task_group.hpp +++ /dev/null @@ -1,51 +0,0 @@ -#pragma once - -#include "kmm/core/domain.hpp" -#include "kmm/core/identifiers.hpp" - -namespace kmm { - -class Runtime; -class TaskGraph; - -struct TaskGroupInit { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskGroupInit) - - public: - Runtime& runtime; - const Domain& domain; -}; - -struct TaskInstance { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskInstance) - - public: - Runtime& runtime; - TaskGraph& graph; - DomainChunk chunk; - MemoryId memory_id; - std::vector buffers; - EventList dependencies; - - size_t add_buffer_requirement(BufferRequirement req) { - size_t index = buffers.size(); - buffers.push_back(std::move(req)); - return index; - } -}; - -struct TaskSubmissionResult { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskSubmissionResult) - - public: - Runtime& runtime; - TaskGraph& graph; - EventId event_id; -}; - -struct TaskGroupCommit { - Runtime& runtime; - TaskGraph& graph; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/api/view_argument.hpp b/include/kmm/api/view_argument.hpp deleted file mode 100644 index 3f4c67cb..00000000 --- a/include/kmm/api/view_argument.hpp +++ /dev/null @@ -1,67 +0,0 @@ -#pragma once - -#include "kmm/api/argument.hpp" -#include "kmm/core/view.hpp" - -namespace kmm { - -template -struct ViewArgument { - using value_type = T; - using domain_type = D; - using mapping_type = typename L::template mapping_type; - - ViewArgument(size_t buffer_index, domain_type domain, mapping_type layout) : - buffer_index(buffer_index), - domain(domain), - layout(layout) {} - - ViewArgument(size_t buffer_index, domain_type domain) : - ViewArgument(buffer_index, domain, L::from_domain(domain)) {} - - size_t buffer_index; - domain_type domain; - mapping_type layout; -}; - -template -struct ArgumentUnpack> { - using type = AbstractView; - - static type call(const TaskContext& context, ViewArgument arg) { - T* data = static_cast(context.accessors.at(arg.buffer_index).address); - return {data, arg.domain, arg.layout}; - } -}; - -template -struct ArgumentUnpack> { - using type = AbstractView; - - static type call(const TaskContext& context, ViewArgument arg) { - const T* data = static_cast(context.accessors.at(arg.buffer_index).address); - return {data, arg.domain, arg.layout}; - } -}; - -template -struct ArgumentUnpack> { - using type = AbstractView; - - static type call(const TaskContext& context, ViewArgument arg) { - T* data = static_cast(context.accessors.at(arg.buffer_index).address); - return {data, arg.domain, arg.layout}; - } -}; - -template -struct ArgumentUnpack> { - using type = AbstractView; - - static type call(const TaskContext& context, ViewArgument arg) { - const T* data = static_cast(context.accessors.at(arg.buffer_index).address); - return {data, arg.domain, arg.layout}; - } -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/core/backends.hpp b/include/kmm/core/backends.hpp deleted file mode 100644 index cb86cc8f..00000000 --- a/include/kmm/core/backends.hpp +++ /dev/null @@ -1,9 +0,0 @@ -#pragma once - -#ifdef KMM_USE_CUDA - #include "kmm/backends/cuda.hpp" -#elif KMM_USE_HIP - #include "kmm/backends/hip.hpp" -#else - #include "kmm/backends/cpu.hpp" -#endif diff --git a/include/kmm/core/bounds.hpp b/include/kmm/core/bounds.hpp new file mode 100644 index 00000000..2aa9a6e4 --- /dev/null +++ b/include/kmm/core/bounds.hpp @@ -0,0 +1,382 @@ +#pragma once + +#include "kmm/core/domain_traits.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/range.hpp" +#include "kmm/core/shape.hpp" +#include "kmm/core/type_utils.hpp" +#include "kmm/core/vec.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// An N-dimensional axis-aligned box, given by a `Range` per axis, backed by a +/// `Vec, N>`. +/// +/// A `Bounds` describes a sub-region of an N-dimensional domain: the half-open `[begin, end)` +/// range of valid indices along each axis. It can be constructed from a begin/end point pair, +/// from an offset and a `Shape`, or directly from a `Shape` (assuming a zero origin). Use +/// `contains`/`overlaps`/`intersection` to test and combine bounds. +template +class Bounds: public Vec, N> { + public: + using storage_type = Vec, N>; + + KMM_HOST_DEVICE + explicit constexpr Bounds(const storage_type& storage) : storage_type(storage) {} + + KMM_HOST_DEVICE + constexpr Bounds() : storage_type(fill(Range())) {} + + constexpr Bounds(const Bounds&) = default; + constexpr Bounds(Bounds&&) noexcept = default; + Bounds& operator=(const Bounds&) = default; + Bounds& operator=(Bounds&&) noexcept = default; + + template> + KMM_HOST_DEVICE Bounds(Range first, Ts&&... args) : storage_type {first, args...} {} + + template + KMM_HOST_DEVICE constexpr Bounds(const Bounds& that) { + if (!that.template is_convertible_to()) { + throw_overflow_exception(); + } + + *this = Bounds::from(that); + } + + /// Construct from begin and end point. + KMM_HOST_DEVICE static constexpr Bounds from_bounds( + const Point& begin, + const Point& end + ) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = {begin[i], end[i]}; + } + + return Bounds(result); + } + + /// Construct from offset and shape. + KMM_HOST_DEVICE static constexpr Bounds from_offset_size( + const Point& offset, + const Shape& shape + ) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = Range(shape[i]) + offset[i]; + } + + return Bounds(result); + } + + KMM_HOST_DEVICE + constexpr Bounds(const Shape& shape) : + Bounds(from_offset_size(Point::zero(), shape)) {} + + /// Returns an empty bounds (all ranges are 0...0). + KMM_HOST_DEVICE static constexpr Bounds empty() { + return Bounds(fill(Range())); + } + + /// Returns `Bounds` with one element (all ranges are 0...1). + KMM_HOST_DEVICE static constexpr Bounds one() { + return Bounds(fill(Range::one())); + } + + /// Returns `Bounds` constructed from another `Bounds`. + template + KMM_HOST_DEVICE static constexpr Bounds from(const Vec, M>& that) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = is_less(i, M) ? Range::from(that[i]) : Range(static_cast(1)); + } + + return Bounds(result); + } + + /// Returns `true` if this `Bounds` can also be represented as `Bounds` without + /// loss of information, `false` otherwise. + template + KMM_HOST_DEVICE bool is_convertible_to() const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + if (i < M) { + result &= (*this)[i].template is_convertible_to(); + } else { + result &= (*this)[i] == Range::one(); + } + } + + return result; + } + + /// Returns the range along the `i`-th dimension and `default_value` if it is out of bounds. + KMM_HOST_DEVICE + Range get_or_default(size_t i, Range default_value = Range::one()) const { + if constexpr (N > 0) { + if (KMM_LIKELY(i < N)) { + return (*this)[i]; + } + } + + return default_value; + } + + /// Returns the begin value along the `i`-th dimension. + KMM_HOST_DEVICE + T begin(size_t axis) const { + return get_or_default(axis).start; + } + + /// Returns the end value along the `i`-th dimension (exclusive). + KMM_HOST_DEVICE + T end(size_t axis) const { + return get_or_default(axis).stop; + } + + /// Returns the number of elements along the `i`-th dimension. + KMM_HOST_DEVICE + T size(size_t axis) const { + return get_or_default(axis).size(); + } + + /// Returns N-d begin point. + KMM_HOST_DEVICE + Point begin() const { + Point result; + for (size_t axis = 0; is_less(axis, N); axis++) { + result[axis] = (*this)[axis].start; + } + return result; + } + + /// Returns N-d end point (exclusive) + KMM_HOST_DEVICE + Point end() const { + Point result; + for (size_t axis = 0; is_less(axis, N); axis++) { + result[axis] = (*this)[axis].stop; + } + return result; + } + + /// Returns the N-d shape (i.e., size along each dimension). + KMM_HOST_DEVICE + Shape shape() const { + Shape result; + for (size_t axis = 0; is_less(axis, N); axis++) { + result[axis] = (*this)[axis].size(); + } + return result; + } + + /// Returns `true` only if this bounds is empty. + KMM_HOST_DEVICE + bool is_empty() const { + bool result = false; + + for (size_t i = 0; is_less(i, N); i++) { + result |= this->begin(i) >= this->end(i); + } + + return result; + } + + /// Returns the product of `size(0) * size(1) * ...`. + KMM_HOST_DEVICE + T volume() const { + T result = 1; + + for (size_t i = 0; is_less(i, N); i++) { + result *= this->end(i) - this->begin(i); + } + + return this->is_empty() ? T {0} : result; + } + + KMM_HOST_DEVICE + Bounds intersection(const Bounds& that) const { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i].start = this->begin(i) >= that.begin(i) ? this->begin(i) : that.begin(i); + result[i].stop = this->end(i) <= that.end(i) ? this->end(i) : that.end(i); + } + + return Bounds(result); + } + + /// Returns `true` if this bounds overlaps the given bounds. + KMM_HOST_DEVICE + bool overlaps(const Bounds& that) const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + result &= this->begin(i) < this->end(i) && that.begin(i) < that.end(i) && // + this->begin(i) < that.end(i) && that.begin(i) < this->end(i); + } + + return result; + } + + /// Returns `true` if this bounds contains the given bounds. + KMM_HOST_DEVICE + bool contains(const Bounds& that) const { + bool is_contained = true; + bool that_is_empty = false; + + for (size_t i = 0; is_less(i, N); i++) { + is_contained &= that.begin(i) >= this->begin(i); + is_contained &= that.end(i) <= this->end(i); + that_is_empty |= that.begin(i) >= that.end(i); + } + + return is_contained || that_is_empty; + } + + /// Returns `true` if this bounds contains the given point. + KMM_HOST_DEVICE + bool contains(const Point& that) const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + result &= (*this)[i].contains(that[i]); + } + + return result; + } + + /// Returns `true` if this bounds contains the given point `{first, rest, ...}`. + // + // NB: uses `enable_if_t<...>` directly rather than the `assert_arity_t` alias -- see the + // comment on the analogous `Shape` constructor for why. + template> + KMM_HOST_DEVICE bool contains(const T& first, Ts&&... rest) const { + return contains(Point {first, rest...}); + } + + /// Returns `true` if this bounds overlaps the given shape. + KMM_HOST_DEVICE + bool overlaps(const Shape& that) const { + return overlaps(Bounds {that}); + } + + /// Returns `true` if this bounds contains the given shape. + KMM_HOST_DEVICE + bool contains(const Shape& that) const { + return contains(Bounds {that}); + } +}; + +template +Bounds(Ts&&...) -> Bounds; + +/// Constructs a `Bounds` from the given ranges. +template +KMM_HOST_DEVICE Bounds bounds(const Ts&... values) { + return Bounds {Vec, sizeof...(Ts)> {values...}}; +} + +/// @} + +namespace detail { + +template +struct domain_traits> { + static constexpr size_t rank = N; + using index_type = IndexT; + using domain_type = Bounds; + template + using drop_axis_type = Bounds; + + KMM_HOST_DEVICE + static constexpr Range bounds(const domain_type& domain, size_t axis) { + return domain[axis]; + } + + KMM_HOST_DEVICE + static constexpr index_type extent(const domain_type& domain, size_t axis) { + return domain[axis].size(); + } + + template + using slice_axis_type = Bounds; + + // Shifts the axis's range by `-begin` and truncates it to length `end - begin`, so the + // returned domain is locally zero-based at this axis; the caller is responsible for + // folding the absolute shift into a storage offset. + template + KMM_HOST_DEVICE static constexpr slice_axis_type slice_axis( + domain_type domain, + index_type begin, + index_type end + ) { + domain[Axis] = Range {end - begin}; + return domain; + } + + template + KMM_HOST_DEVICE static constexpr drop_axis_type drop_axis(const domain_type& domain) { + return permute_axes(domain, drop_index_sequence()); + } + + template + using permute_axes_type = Bounds; + + template + KMM_HOST_DEVICE static constexpr permute_axes_type permute_axes( + const domain_type& domain, + IndexSequence + ) { + return {domain[Is]...}; + } + + template + using insert_axis_type = Bounds; + + template + KMM_HOST_DEVICE static constexpr insert_axis_type insert_axis( + const domain_type& domain, + index_type extent + ) { + insert_axis_type result; + + for (size_t i = 0; is_less(i, Axis); i++) { + result[i] = domain[i]; + } + + result[Axis] = Range(extent); + + for (size_t i = Axis; is_less(i, N); i++) { + result[i + 1] = domain[i]; + } + + return result; + } +}; + +} // namespace detail + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash>: std::hash, N>> {}; +#endif \ No newline at end of file diff --git a/include/kmm/core/buffer.hpp b/include/kmm/core/buffer.hpp deleted file mode 100644 index 77b8e1b2..00000000 --- a/include/kmm/core/buffer.hpp +++ /dev/null @@ -1,79 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "kmm/core/data_type.hpp" -#include "kmm/core/identifiers.hpp" - -namespace kmm { - -/** - * Represents the layout of a buffer. For now, this is just its size and alignment. - */ -struct BufferLayout { - BufferLayout repeat(size_t n) { - size_t remainder = size_in_bytes % alignment; - size_t padding = remainder != 0 ? alignment - remainder : 0; - return {(size_in_bytes + padding) * n, alignment}; - } - - template - static BufferLayout for_type(size_t n = 1) { - return BufferLayout {sizeof(T), alignof(T)}.repeat(n); - } - - static BufferLayout for_type(DataType dtype, size_t n = 1) { - return BufferLayout {dtype.size_in_bytes(), dtype.alignment()}.repeat(n); - } - - size_t size_in_bytes = 0; - size_t alignment = 1; -}; - -/** - * This enum is used to specify how a buffer can be accessed: read-only, read-write, or exclusive. - */ -enum struct AccessMode { - Read, ///< Read-only access to the buffer. - ReadWrite, ///< Read and write access to the buffer. - Exclusive ///< Exclusive access, implying full control over the buffer. -}; - -/** - * Represents the requirements for accessing a buffer. - */ -struct BufferRequirement { - BufferId buffer_id; - MemoryId memory_id; - AccessMode access_mode; -}; - -/** - * Provides access to a buffer with specific properties. - */ -struct BufferAccessor { - MemoryId memory_id; - BufferLayout layout; - bool is_writable; - void* address; -}; - -inline std::ostream& operator<<(std::ostream& f, AccessMode mode) { - switch (mode) { - case AccessMode::Read: - return f << "Read"; - case AccessMode::ReadWrite: - return f << "ReadWrite"; - case AccessMode::Exclusive: - return f << "Exclusive"; - } - - return f; -} - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/core/checked_compare.hpp b/include/kmm/core/checked_compare.hpp new file mode 100644 index 00000000..3c83e0be --- /dev/null +++ b/include/kmm/core/checked_compare.hpp @@ -0,0 +1,571 @@ +#pragma once + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" + +namespace kmm { + +namespace detail { + +enum class numeric_type_tag { // + signed_int, + unsigned_int, + floating_point, + other +}; + +template +struct numeric_type_traits { + static constexpr numeric_type_tag tag = numeric_type_tag::other; +}; + +template<> +struct numeric_type_traits { + static constexpr numeric_type_tag tag = numeric_type_tag::floating_point; + + KMM_HOST_DEVICE + static constexpr bool isnan(float f) { + return f != f; + } + + KMM_HOST_DEVICE + static float ceil(float f) { + return ::ceilf(f); + } + + KMM_HOST_DEVICE + static float floor(float f) { + return ::floorf(f); + } +}; + +template<> +struct numeric_type_traits { + static constexpr numeric_type_tag tag = numeric_type_tag::floating_point; + + KMM_HOST_DEVICE + static constexpr bool isnan(double f) { + return f != f; + } + + KMM_HOST_DEVICE + static double ceil(double f) { + return ::ceil(f); + } + + KMM_HOST_DEVICE + static double floor(double f) { + return ::floor(f); + } +}; + +template<> +struct numeric_type_traits { + static constexpr bool is_signed = false; + using unsigned_type = bool; + + static constexpr numeric_type_tag tag = numeric_type_tag::unsigned_int; + static constexpr bool min_inclusive = false; + static constexpr bool max_inclusive = true; + + static constexpr float min_inclusive_float = 0.0F; + static constexpr float max_exclusive_float = 2.0F; +}; + +#define KMM_DEFINE_INT_TRAITS(T) \ + template<> \ + struct numeric_type_traits { \ + static constexpr bool is_signed = true; \ + static constexpr numeric_type_tag tag = numeric_type_tag::signed_int; \ + using unsigned_type = unsigned T; \ + \ + static constexpr signed T max_inclusive = \ + (signed T)(unsigned_type(~unsigned_type(0)) >> 1); \ + static constexpr signed T min_inclusive = ~max_inclusive; \ + \ + static constexpr float min_inclusive_float = min_inclusive; \ + static constexpr float max_exclusive_float = unsigned_type(max_inclusive) + 1; \ + }; \ + \ + template<> \ + struct numeric_type_traits { \ + static constexpr bool is_signed = false; \ + static constexpr numeric_type_tag tag = numeric_type_tag::unsigned_int; \ + using unsigned_type = unsigned T; \ + \ + static constexpr unsigned T min_inclusive = 0; \ + static constexpr unsigned T max_inclusive = ~static_cast(0); \ + \ + static constexpr float min_inclusive_float = min_inclusive; \ + static constexpr float max_exclusive_float = 2.0f * float(max_inclusive / 2 + 1); \ + }; + +KMM_DEFINE_INT_TRAITS(char) +KMM_DEFINE_INT_TRAITS(short) +KMM_DEFINE_INT_TRAITS(int) +KMM_DEFINE_INT_TRAITS(long) +KMM_DEFINE_INT_TRAITS(long long) + +template +struct checked_compare_base { + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + return left == right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + return left < right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return left <= right; + } +}; + +template< + typename L, + typename R, + numeric_type_tag = numeric_type_traits::tag, + numeric_type_tag = numeric_type_traits::tag> +struct checked_compare_impl; + +template +struct checked_compare_impl: + checked_compare_base {}; + +template +struct checked_compare_impl: + checked_compare_base {}; + +template +struct checked_compare_impl: + checked_compare_base {}; + +template +struct checked_compare_impl { + using UR = typename numeric_type_traits::unsigned_type; + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + return right >= static_cast(0) && left == static_cast(right); + } + + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + return right >= static_cast(0) && left < static_cast(right); + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return right >= static_cast(0) && left <= static_cast(right); + } +}; + +template +struct checked_compare_impl { + using UL = typename numeric_type_traits::unsigned_type; + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + return left >= static_cast(0) && static_cast

    (left) == right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + return left < static_cast(0) || static_cast
      (left) < right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return left < static_cast(0) || static_cast
        (left) <= right; + } +}; + +template +struct checked_compare_impl< + L, + R, + numeric_type_tag::floating_point, + numeric_type_tag::floating_point>: checked_compare_base {}; + +template +struct checked_compare_impl { + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + if (numeric_type_traits::isnan(left)) { + return false; + } + + if (numeric_type_traits::floor(left) < numeric_type_traits::min_inclusive_float) { + return true; + } + + if (!(numeric_type_traits::floor(left) < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(numeric_type_traits::floor(left)) < right; + } + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + if (numeric_type_traits::floor(left) != left) { + return false; + } + + if (left < numeric_type_traits::min_inclusive_float) { + return false; + } + + if (!(left < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(left) == right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return is_less(left, right) || is_equal(left, right); + } +}; + +template +struct checked_compare_impl { + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + if (numeric_type_traits::isnan(right)) { + return false; + } + + if (numeric_type_traits::ceil(right) < numeric_type_traits::min_inclusive_float) { + return false; + } + + if (!(numeric_type_traits::ceil(right) < numeric_type_traits::max_exclusive_float)) { + return true; + } + + return left < static_cast(numeric_type_traits::ceil(right)); + } + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + return checked_compare_impl::is_equal(right, left); + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return is_less(left, right) || is_equal(left, right); + } +}; + +template +struct checked_compare_impl< + L, + R, + numeric_type_tag::floating_point, + numeric_type_tag::unsigned_int> { + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + if (numeric_type_traits::isnan(left)) { + return false; + } + + if (numeric_type_traits::floor(left) < numeric_type_traits::min_inclusive_float) { + return true; + } + + if (!(numeric_type_traits::floor(left) < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(numeric_type_traits::floor(left)) < right; + } + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + if (numeric_type_traits::isnan(left)) { + return false; + } + + if (numeric_type_traits::floor(left) != left) { + return false; + } + + if (left < numeric_type_traits::min_inclusive_float) { + return false; + } + + if (!(left < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(left) == right; + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return is_less(left, right) || is_equal(left, right); + } +}; + +template +struct checked_compare_impl< + L, + R, + numeric_type_tag::unsigned_int, + numeric_type_tag::floating_point> { + KMM_HOST_DEVICE + static constexpr bool is_less(L left, R right) { + if (numeric_type_traits::isnan(right)) { + return false; + } + + if (numeric_type_traits::ceil(right) < numeric_type_traits::min_inclusive_float) { + return false; + } + + if (!(numeric_type_traits::ceil(right) < numeric_type_traits::max_exclusive_float)) { + return true; + } + + return left < static_cast(numeric_type_traits::ceil(right)); + } + + KMM_HOST_DEVICE + static constexpr bool is_equal(L left, R right) { + return checked_compare_impl::is_equal(right, left); + } + + KMM_HOST_DEVICE + static constexpr bool is_less_equal(L left, R right) { + return is_less(left, right) || is_equal(left, right); + } +}; + +template< + typename I, + typename O, + numeric_type_tag = numeric_type_traits::tag, + numeric_type_tag = numeric_type_traits::tag> +struct is_convertible_impl; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const T& input) { + return true; + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + return input >= numeric_type_traits::min_inclusive + && input <= numeric_type_traits::max_inclusive; + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + return input <= numeric_type_traits::max_inclusive; + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + using UO = typename numeric_type_traits::unsigned_type; + return input <= static_cast(numeric_type_traits::max_inclusive); + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + using UI = typename numeric_type_traits::unsigned_type; + return input >= static_cast(0) + && is_convertible_impl::apply(static_cast(input)); + } +}; + +template +struct is_convertible_impl< + I, + O, + numeric_type_tag::floating_point, + numeric_type_tag::floating_point> { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + return input == static_cast(input); + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + if (input < numeric_type_traits::min_inclusive_float) { + return false; + } + + if (!(input < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(static_cast(input)) == input; + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + if (input < static_cast(0)) { + return false; + } + + if (!(input < numeric_type_traits::max_exclusive_float)) { + return false; + } + + return static_cast(static_cast(input)) == input; + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + return checked_compare_impl::is_equal(input, static_cast(input)); + } +}; + +template +struct is_convertible_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input) { + return checked_compare_impl::is_equal(input, static_cast(input)); + } +}; + +template +struct checked_cast_impl { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input, O* output) { + if (!is_convertible_impl::apply(input)) { + return false; + } + + *output = static_cast(input); + return true; + } +}; +} // namespace detail + +#if KMM_IS_DEVICE +// on the GPU, we just panic immediately +KMM_DEVICE void throw_overflow_exception() { + KMM_PANIC("overflow occurred in operation"); +} +#else +// on the host, we can throw an exception +[[noreturn]] void throw_overflow_exception(); +#endif + +/// \addtogroup utility +/// @{ + +/** + * All functions below (`is_less`, `is_equal`, etc.) cast their arguments to the underlying types + * before performing the operation. By default, the underlying type of `T` is just `T`. + * + * However, other types can override this by specializing `underlying_type` for their own + * type and setting `type` to whatever native type it should be compared/converted as. + */ +template +struct underlying_type { + using type = T; +}; + +template +using underlying_type_t = typename underlying_type::type; + +/// Returns whether `left < right`, safely comparing operands of different types and signedness. +template +KMM_HOST_DEVICE constexpr bool is_less(L left, R right) { + return detail::checked_compare_impl, underlying_type_t>::is_less( + static_cast>(left), + static_cast>(right) + ); +} + +/// Returns whether `left == right`, safely comparing operands of different types and signedness. +template +KMM_HOST_DEVICE constexpr bool is_equal(L left, R right) { + return detail::checked_compare_impl, underlying_type_t>::is_equal( + static_cast>(left), + static_cast>(right) + ); +} + +/// Returns whether `left <= right`, safely comparing operands of different types and signedness. +template +KMM_HOST_DEVICE constexpr bool is_less_equal(L left, R right) { + return detail::checked_compare_impl, underlying_type_t>::is_less_equal( + static_cast>(left), + static_cast>(right) + ); +} + +/// Returns whether `left > right`, safely comparing operands of different types and signedness. +template +KMM_HOST_DEVICE constexpr bool is_greater(L left, R right) { + return is_less(right, left); +} + +/// Returns whether `left >= right`, safely comparing operands of different types and signedness. +template +KMM_HOST_DEVICE constexpr bool is_greater_equal(L left, R right) { + return is_less_equal(right, left); +} + +/// Returns `true` if the given `input` value can safely be converted from type `T` to type `U`, +/// and `false` otherwise. For example, `int32_t(100)` can be converted to `uint16_t`, but +/// `int32_t(-100)` cannot since it cannot be represented as a `uint16_t`. +template +KMM_HOST_DEVICE constexpr bool is_convertible(const T& input) { + U output {}; + return detail::checked_cast_impl, U>::apply( + static_cast>(input), + &output + ); +} + +/// Casts the given value `input` from type `T` to type `U`. Throws an exception if the +/// input value cannot be represented as `U`. +template +KMM_HOST_DEVICE constexpr U checked_cast(const T& input) { + U output {}; + + if (!detail::checked_cast_impl, U>::apply( + static_cast>(input), + &output + )) { + throw_overflow_exception(); + } + + return output; +} + +/// @} + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/core/checked_math.hpp b/include/kmm/core/checked_math.hpp new file mode 100644 index 00000000..bc77c723 --- /dev/null +++ b/include/kmm/core/checked_math.hpp @@ -0,0 +1,431 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/macros.hpp" + +// __mul64hi/__umul64hi are bare compiler builtins under CUDA, but HIP only declares them once its +// runtime header is included. +#if KMM_IS_DEVICE + #if defined(__CUDACC__) + #include + #elif defined(__HIPCC__) + #include + #endif +#endif + +namespace kmm { + +namespace detail { + +template::tag> +struct wider_type; + +template +struct wider_type { + using type = int64_t; +}; + +template +struct wider_type { + using type = uint64_t; +}; + +template +using wider_t = typename wider_type::type; + +// `value < 0` guarded so it never becomes a "comparison of unsigned with zero" for unsigned `T`. +template +KMM_HOST_DEVICE constexpr bool is_negative(T value) { + if constexpr (numeric_type_traits::tag == numeric_type_tag::signed_int) { + return value < static_cast(0); + } else { + return false; + } +} + +template +struct checked_add_impl { + KMM_HOST_DEVICE + static constexpr bool apply(L left, R right, O* output) { + uint64_t sum = static_cast(static_cast>(left)) + + static_cast(static_cast>(right)); + *output = static_cast(sum); + + bool left_negative = is_negative(left); + bool right_negative = is_negative(right); + bool carry = sum < static_cast(static_cast>(left)); + bool sum_negative = numeric_type_traits::is_signed && (static_cast(sum) < 0); + + bool is_valid = int(carry) - int(left_negative) - int(right_negative) == -int(sum_negative); + + // if all are signed, the above can be simplified to just this: + // same-sign operands overflow iff the result's sign differs from both. + if constexpr (numeric_type_traits::is_signed && // + numeric_type_traits::is_signed && // + numeric_type_traits::is_signed) { + int64_t l = static_cast(left); + int64_t r = static_cast(right); + int64_t s = static_cast(sum); + is_valid = ((s ^ l) & (s ^ r)) >= 0; + } + + return is_valid && is_convertible_impl, O>::apply(static_cast>(sum)); + } +}; + +template +struct checked_sub_impl { + KMM_HOST_DEVICE + static constexpr bool apply(L left, R right, O* output) { + uint64_t diff = static_cast(static_cast>(left)) + - static_cast(static_cast>(right)); + *output = static_cast(diff); + + bool left_negative = is_negative(left); + bool right_negative = is_negative(right); + bool borrow = static_cast(static_cast>(left)) + < static_cast(static_cast>(right)); + bool diff_negative = numeric_type_traits::is_signed && static_cast(diff) < 0; + + bool is_valid = + int(borrow) - int(right_negative) - int(diff_negative) == -int(left_negative); + + // if all are signed, the above can be simplified to just this: + // same-sign operands never overflow; for mixed signs, overflow iff the result's + // sign differs from `left`. + if constexpr (numeric_type_traits::is_signed && // + numeric_type_traits::is_signed && // + numeric_type_traits::is_signed) { + int64_t l = static_cast(left); + int64_t r = static_cast(right); + int64_t d = static_cast(diff); + is_valid = ((l ^ r) & (l ^ d)) >= 0; + } + + return is_valid && is_convertible_impl, O>::apply(static_cast>(diff)); + } +}; + +template +struct checked_mul_impl { + KMM_HOST_DEVICE + static constexpr bool apply(L left, R right, O* output) { +#if KMM_IS_DEVICE + // Fast path when every operand is signed + if constexpr (numeric_type_traits::is_signed && // + numeric_type_traits::is_signed && // + numeric_type_traits::is_signed) { + int64_t l = static_cast(left); + int64_t r = static_cast(right); + int64_t lo = static_cast(static_cast(l) * static_cast(r)); + + // must have same sign. + if (__mul64hi(l, r) != (lo >> 63)) { + return false; + } + + *output = static_cast(lo); + return is_convertible_impl::apply(lo); + } + + // General path (mixed signedness): reduce both operands to unsigned magnitudes, use the + // hardware high-multiply to detect a product wider than 64 bits, then reapply the sign. + uint64_t al = static_cast(static_cast>(left)); + uint64_t ar = static_cast(static_cast>(right)); + + if (is_negative(left)) { + al = -al; + } + + if (is_negative(right)) { + ar = -ar; + } + + if (__umul64hi(al, ar) != 0) { + return false; + } + + uint64_t magnitude = al * ar; + bool result_negative = is_negative(left) != is_negative(right); + + if (result_negative) { + return checked_sub_impl::apply(uint64_t(0), magnitude, output); + } else { + return checked_add_impl::apply(uint64_t(0), magnitude, output); + } +#else + return !__builtin_mul_overflow(left, right, output); +#endif + } +}; + +template +struct checked_div_impl { + KMM_HOST_DEVICE + static constexpr bool apply(L left, R right, O* output) { + // division by zero + if (right == static_cast(0)) { + return false; + } + + // If both are signed, the result fits into int64_t, except for MIN/-1 + if constexpr (numeric_type_traits::is_signed && numeric_type_traits::is_signed) { + if (left == numeric_type_traits::min_inclusive && right == -1) { + uint64_t magnitude = uint64_t(numeric_type_traits::max_inclusive) + 1; + *output = static_cast(magnitude); + return is_convertible_impl::apply(magnitude); + } else { + int64_t magnitude = static_cast(left) / static_cast(right); + *output = static_cast(magnitude); + return is_convertible_impl::apply(magnitude); + } + } + + uint64_t al = static_cast(static_cast>(left)); + uint64_t ar = static_cast(static_cast>(right)); + + if (is_negative(left)) { + al = -al; + } + + if (is_negative(right)) { + ar = -ar; + } + + uint64_t magnitude = al / ar; + bool result_negative = is_negative(left) != is_negative(right); + + if (result_negative) { + return checked_sub_impl::apply(uint64_t(0), magnitude, output); + } else { + return checked_add_impl::apply(uint64_t(0), magnitude, output); + } + } +}; + +template +struct checked_rem_impl { + KMM_HOST_DEVICE + static constexpr bool apply(L left, R right, O* output) { + // division by zero + if (right == static_cast(0)) { + return false; + } + + // If both are signed, the result fits into int64_t, except for MIN/-1 + if constexpr (numeric_type_traits::is_signed && numeric_type_traits::is_signed) { + int64_t magnitude; + + if (left == numeric_type_traits::min_inclusive && right == -1) { + magnitude = static_cast(0); + } else { + magnitude = static_cast(left) % static_cast(right); + } + + *output = static_cast(magnitude); + return is_convertible_impl::apply(magnitude); + } + + uint64_t al = static_cast(static_cast>(left)); + uint64_t ar = static_cast(static_cast>(right)); + + if (is_negative(left)) { + al = -al; + } + + if (is_negative(right)) { + ar = -ar; + } + + uint64_t magnitude = al % ar; + bool result_negative = is_negative(left); // sign follows dividend + + if (result_negative) { + return checked_sub_impl::apply(uint64_t(0), magnitude, output); + } else { + return checked_add_impl::apply(uint64_t(0), magnitude, output); + } + } +}; + +// These `__builtin_*_overflow` compiler builtins have no device-side implementation, so these +// specializations are compiled for the host pass only; the device pass never sees them declared +// and falls back to the generic, non-specialized `checked_*_impl::apply` above instead. +#if !KMM_IS_DEVICE + #define KMM_CHECKED_MATH_IMPL(OP, T, FUN) \ + template<> \ + struct OP { \ + KMM_HOST_DEVICE \ + static bool apply(T left, T right, T* output) { \ + return FUN(left, right, output) == false; \ + } \ + }; + +KMM_CHECKED_MATH_IMPL(checked_add_impl, int, __builtin_sadd_overflow) +KMM_CHECKED_MATH_IMPL(checked_add_impl, long, __builtin_saddl_overflow) +KMM_CHECKED_MATH_IMPL(checked_add_impl, long long, __builtin_saddll_overflow) +KMM_CHECKED_MATH_IMPL(checked_add_impl, unsigned int, __builtin_uadd_overflow) +KMM_CHECKED_MATH_IMPL(checked_add_impl, unsigned long, __builtin_uaddl_overflow) +KMM_CHECKED_MATH_IMPL(checked_add_impl, unsigned long long, __builtin_uaddll_overflow) + +KMM_CHECKED_MATH_IMPL(checked_sub_impl, int, __builtin_ssub_overflow) +KMM_CHECKED_MATH_IMPL(checked_sub_impl, long, __builtin_ssubl_overflow) +KMM_CHECKED_MATH_IMPL(checked_sub_impl, long long, __builtin_ssubll_overflow) +KMM_CHECKED_MATH_IMPL(checked_sub_impl, unsigned int, __builtin_usub_overflow) +KMM_CHECKED_MATH_IMPL(checked_sub_impl, unsigned long, __builtin_usubl_overflow) +KMM_CHECKED_MATH_IMPL(checked_sub_impl, unsigned long long, __builtin_usubll_overflow) + +KMM_CHECKED_MATH_IMPL(checked_mul_impl, int, __builtin_smul_overflow) +KMM_CHECKED_MATH_IMPL(checked_mul_impl, long, __builtin_smull_overflow) +KMM_CHECKED_MATH_IMPL(checked_mul_impl, long long, __builtin_smulll_overflow) +KMM_CHECKED_MATH_IMPL(checked_mul_impl, unsigned int, __builtin_umul_overflow) +KMM_CHECKED_MATH_IMPL(checked_mul_impl, unsigned long, __builtin_umull_overflow) +KMM_CHECKED_MATH_IMPL(checked_mul_impl, unsigned long long, __builtin_umulll_overflow) + #undef KMM_CHECKED_MATH_IMPL +#endif // !KMM_IS_DEVICE + +} // namespace detail + +/// \addtogroup utility +/// @{ + +/// Returns `left + right`, throwing on overflow. +template +KMM_HOST_DEVICE constexpr O checked_add(L left, R right) { + O output {}; + + if (!detail::checked_add_impl::apply(left, right, &output)) { + throw_overflow_exception(); + } + + return output; +} + +template +KMM_HOST_DEVICE constexpr T checked_add(T left, T right) { + return checked_add(left, right); +} + +/// Returns `left - right`, throwing on overflow. +template +KMM_HOST_DEVICE constexpr O checked_sub(L left, R right) { + O output {}; + + if (!detail::checked_sub_impl::apply(left, right, &output)) { + throw_overflow_exception(); + } + + return output; +} + +template +KMM_HOST_DEVICE constexpr T checked_sub(T left, T right) { + return checked_sub(left, right); +} + +/// Returns `left * right`, throwing on overflow. +template +KMM_HOST_DEVICE constexpr O checked_mul(L left, R right) { + O output {}; + + if (!detail::checked_mul_impl::apply(left, right, &output)) { + throw_overflow_exception(); + } + + return output; +} + +template +KMM_HOST_DEVICE constexpr T checked_mul(T left, T right) { + return checked_mul(left, right); +} + +/// Returns `left / right`, throwing on division by zero or overflow. +template +KMM_HOST_DEVICE constexpr O checked_div(L left, R right) { + O output {}; + + if (!detail::checked_div_impl::apply(left, right, &output)) { + throw_overflow_exception(); + } + + return output; +} + +template +KMM_HOST_DEVICE constexpr T checked_div(T left, T right) { + return checked_div(left, right); +} + +/// Returns `left % right`, throwing on division by zero or overflow. +template +KMM_HOST_DEVICE constexpr O checked_rem(L left, R right) { + O output {}; + + if (!detail::checked_rem_impl::apply(left, right, &output)) { + throw_overflow_exception(); + } + + return output; +} + +template +KMM_HOST_DEVICE constexpr T checked_rem(T left, T right) { + return checked_rem(left, right); +} + +/// Returns `-input`, throwing on overflow. +template +KMM_HOST_DEVICE constexpr T checked_neg(const T& input) { + return checked_sub(static_cast(0), input); +} + +/// Returns `abs(input)`, throwing on overflow. +template +KMM_HOST_DEVICE constexpr T checked_abs(const T& input) { + return is_less(input, static_cast(0)) ? checked_neg(input) : input; +} + +/// Returns `begin[0] + begin[1] + ... + begin[end - begin]`, throwing on overflow. +template +KMM_HOST_DEVICE U checked_sum(It begin, It end, U initial = U(0)) { + using T = decltype(+*begin); + U accum = initial; + bool is_valid = true; + + for (It it = begin; it != end; ++it) { + if (!detail::checked_add_impl::apply(accum, *it, &accum)) { + is_valid = false; + } + } + + if (!is_valid) { + throw_overflow_exception(); + } + + return accum; +} + +/// Returns `begin[0] * begin[1] * ... * begin[end - begin]`, throwing on overflow. +template +KMM_HOST_DEVICE U checked_product(It begin, It end, U initial = U(1)) { + using T = decltype(+*begin); + bool is_valid = true; + U accum = initial; + + for (It it = begin; it != end; ++it) { + if (!detail::checked_mul_impl::apply(accum, *it, &accum)) { + is_valid = false; + } + } + + if (!is_valid) { + throw_overflow_exception(); + } + + return accum; +} + +/// @} + +} // namespace kmm diff --git a/include/kmm/core/commands.hpp b/include/kmm/core/commands.hpp deleted file mode 100644 index 8c53137e..00000000 --- a/include/kmm/core/commands.hpp +++ /dev/null @@ -1,83 +0,0 @@ -#pragma once - -#include - -#include "fmt/ostream.h" - -#include "kmm/core/buffer.hpp" -#include "kmm/core/identifiers.hpp" -#include "kmm/core/reduction.hpp" -#include "kmm/core/resource.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -struct CommandBufferDelete { - BufferId id; -}; - -struct CommandPrefetch { - BufferId buffer_id; - MemoryId memory_id; -}; - -struct CommandCopy { - BufferId src_buffer; - MemoryId src_memory; - BufferId dst_buffer; - MemoryId dst_memory; - CopyDef definition; -}; - -struct CommandExecute { - ResourceId processor_id; - std::unique_ptr task; - std::vector buffers; -}; - -struct CommandReduction { - BufferId src_buffer; - BufferId dst_buffer; - MemoryId memory_id; - ReductionDef definition; -}; - -struct CommandFill { - BufferId dst_buffer; - MemoryId memory_id; - FillDef definition; -}; - -struct CommandEmpty {}; - -using Command = std::variant< - CommandEmpty, - CommandBufferDelete, - CommandPrefetch, - CommandCopy, - CommandExecute, - CommandReduction, - CommandFill>; - -inline const char* command_name(const Command& cmd) { - static constexpr const char* names[] = { - "CommandEmpty", - "CommandBufferDelete", - "CommandPrefetch", - "CommandCopy", - "CommandExecute", - "CommandReduction", - "CommandFill" - }; - - return names[cmd.index()]; -} - -inline std::ostream& operator<<(std::ostream& f, const Command& cmd) { - return f << command_name(cmd); -} - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; diff --git a/include/kmm/core/const_value.hpp b/include/kmm/core/const_value.hpp new file mode 100644 index 00000000..3439dd58 --- /dev/null +++ b/include/kmm/core/const_value.hpp @@ -0,0 +1,93 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/checked_math.hpp" +#include "kmm/core/macros.hpp" + +namespace kmm { + +namespace detail { +template(OtherValue)> +struct is_convertible_helper {}; + +template +struct is_convertible_helper { + using type = void; +}; +} // namespace detail + +/// Represents a value that is known at compile time. +template +struct ConstValue { + using type = ConstValue; + using value_type = decltype(Value); + static constexpr value_type value = Value; + + KMM_HOST_DEVICE + constexpr ConstValue() = default; + + template< + auto OtherValue, + typename = typename detail::is_convertible_helper::type> + KMM_HOST_DEVICE constexpr ConstValue(ConstValue) noexcept {} + + KMM_HOST_DEVICE + constexpr operator value_type() const noexcept { + return Value; + } + + KMM_HOST_DEVICE + constexpr value_type operator()() const noexcept { + return value; + } +}; + +template +using ConstIndex = ConstValue; + +template +using ConstBool = ConstValue; + +template +ConstValue<+Value> operator+(ConstValue) { + return {}; +} + +template +ConstValue<-Value> operator-(ConstValue) { + checked_neg(Value); + return {}; +} + +#define KMM_CONST_VALUE_OP(OP, FUN) \ + template \ + ConstValue operator OP(ConstValue, ConstValue) { \ + FUN(Left, Right); \ + return {}; \ + } + +KMM_CONST_VALUE_OP(+, checked_add) +KMM_CONST_VALUE_OP(-, checked_sub) +KMM_CONST_VALUE_OP(*, checked_mul) +KMM_CONST_VALUE_OP(/, checked_div) + +#undef KMM_CONST_VALUE_OP + +// ConstValue should compare/convert exactly like a plain `decltype(Value)`. +template +struct underlying_type>: underlying_type {}; + +namespace detail { + +// allows for `checked_cast>(x)`. +template +struct checked_cast_impl> { + KMM_HOST_DEVICE + static constexpr bool apply(const I& input, ConstValue* output) { + return is_equal(input, Value); + } +}; + +} // namespace detail + +} // namespace kmm diff --git a/include/kmm/core/data_type.hpp b/include/kmm/core/data_type.hpp deleted file mode 100644 index 46968971..00000000 --- a/include/kmm/core/data_type.hpp +++ /dev/null @@ -1,120 +0,0 @@ -#pragma once - -#include -#include -#include -#include - -#include "fmt/ostream.h" - -#include "kmm/utils/key_value.hpp" - -namespace kmm { - -enum struct ScalarType : uint8_t { - Invalid = 0, - Int8, - Int16, - Int32, - Int64, - Uint8, - Uint16, - Uint32, - Uint64, - Float16, - Float32, - Float64, - BFloat16, - Complex16, - Complex32, - Complex64, - KeyAndInt64, - KeyAndFloat64, -}; - -template -struct DataTypeOf; - -struct DataType { - DataType() = default; - - static DataType of(ScalarType kind); - - template - static DataType of() { - static_assert(DataTypeOf::value.size_in_bytes == sizeof(T)); - static_assert(DataTypeOf::value.alignment >= alignof(T)); - return DataType {&DataTypeOf::value}; - } - - const std::type_info& type_info() const; - - ScalarType as_scalar() const; - - size_t size_in_bytes() const; - - size_t alignment() const; - - const char* name() const; - - const char* c_name() const; - - public: - struct Info { - size_t size_in_bytes; - size_t alignment; - const char* name = nullptr; - const char* c_name = nullptr; - const std::type_info& type_id; - ScalarType scalar_type = ScalarType::Invalid; - }; - - private: - explicit DataType(const Info* info) : m_info(info) {} - const Info* m_info = nullptr; -}; - -template -struct DataTypeOf { - static constexpr DataType::Info value = - {.size_in_bytes = sizeof(T), .alignment = alignof(T), .type_id = typeid(T)}; -}; - -std::ostream& operator<<(std::ostream& f, ScalarType p); -std::ostream& operator<<(std::ostream& f, DataType p); - -} // namespace kmm - -#define KMM_DEFINE_SCALAR_TYPE(S, T) \ - template<> \ - struct kmm::DataTypeOf { \ - static constexpr ::kmm::DataType::Info value = { \ - .size_in_bytes = sizeof(T), \ - .alignment = alignof(T), \ - .name = #S, \ - .c_name = #T, \ - .type_id = typeid(T), \ - .scalar_type = ::kmm::ScalarType::S \ - }; \ - }; - -KMM_DEFINE_SCALAR_TYPE(Int8, int8_t) -KMM_DEFINE_SCALAR_TYPE(Int16, int16_t) -KMM_DEFINE_SCALAR_TYPE(Int32, int32_t) -KMM_DEFINE_SCALAR_TYPE(Int64, int64_t) -KMM_DEFINE_SCALAR_TYPE(Uint8, uint8_t) -KMM_DEFINE_SCALAR_TYPE(Uint16, uint16_t) -KMM_DEFINE_SCALAR_TYPE(Uint32, uint32_t) -KMM_DEFINE_SCALAR_TYPE(Uint64, uint64_t) -KMM_DEFINE_SCALAR_TYPE(Float32, float) -KMM_DEFINE_SCALAR_TYPE(Float64, double) -KMM_DEFINE_SCALAR_TYPE(Complex32, ::std::complex) -KMM_DEFINE_SCALAR_TYPE(Complex64, ::std::complex) -KMM_DEFINE_SCALAR_TYPE(KeyAndInt64, ::kmm::KeyValue) -KMM_DEFINE_SCALAR_TYPE(KeyAndFloat64, ::kmm::KeyValue) - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/core/distribution.hpp b/include/kmm/core/distribution.hpp index 8388b809..5e31a27c 100644 --- a/include/kmm/core/distribution.hpp +++ b/include/kmm/core/distribution.hpp @@ -1,83 +1,164 @@ #pragma once -#include "kmm/core/buffer.hpp" -#include "kmm/core/domain.hpp" -#include "kmm/core/identifiers.hpp" -#include "kmm/utils/geometry.hpp" +#include "kmm/core/bounds.hpp" +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/shape.hpp" namespace kmm { -template -struct ArrayChunk { - MemoryId owner_id; - Point offset; - Dim size; -}; +/// \addtogroup geometry +/// @{ -template +/// Describes how an N-dimensional `total_shape` is partitioned into a row-major grid of +/// chunks, each of (at most) `chunk_shape` elements along every axis. +template class Distribution { public: - Distribution(); - Distribution(Dim array_size, Dim chunk_size, std::vector memories); + using index_type = IndexT; + using shape_type = Shape; + using point_type = Point; + + KMM_HOST_DEVICE + constexpr Distribution() = default; + + KMM_HOST_DEVICE + constexpr Distribution(shape_type total_shape, shape_type chunk_shape) : + m_total_shape(total_shape), + m_chunk_shape(chunk_shape) { + shape_type grid_shape; + + for (size_t i = 0; is_less(i, N); i++) { + auto extent = total_shape[i]; + auto chunk = chunk_shape[i]; + grid_shape[i] = chunk > static_cast(0) + ? static_cast((extent + chunk - 1) / chunk) + : static_cast(0); + } + + m_grid_shape = grid_shape; + } - static Distribution from_chunks( - Dim array_size, - std::vector> chunks, - bool allow_duplicates = false - ); + /// The full extent of the distributed domain. + KMM_HOST_DEVICE + shape_type total_shape() const noexcept { + return m_total_shape; + } - size_t region_to_chunk_index(Bounds region) const; + /// The nominal (maximum) extent of a single chunk; chunks at the edge of the domain may be + /// smaller, see `chunk_extent`. + KMM_HOST_DEVICE + shape_type chunk_shape() const noexcept { + return m_chunk_shape; + } - ArrayChunk chunk(size_t index) const; + /// The number of chunks along each axis. + KMM_HOST_DEVICE + shape_type grid_shape() const noexcept { + return m_grid_shape; + } - size_t num_chunks() const { - return m_memories.size(); + /// The total number of chunks (`grid_shape().volume()`). + KMM_HOST_DEVICE + size_t num_chunks() const noexcept { + return static_cast(m_grid_shape.volume()); } - Dim chunk_size() const { - return m_chunk_size; + /// The actual extent of the chunk at the given grid index, clipped at the edges of + /// `total_shape`. + KMM_HOST_DEVICE + shape_type chunk_extent(point_type grid_index) const noexcept { + shape_type extent; + + for (size_t i = 0; is_less(i, N); i++) { + auto begin = grid_index[i] * m_chunk_shape[i]; + auto end = begin + m_chunk_shape[i]; + + if (end > m_total_shape[i]) { + end = m_total_shape[i]; + } + + extent[i] = end > begin ? end - begin : static_cast(0); + } + + return extent; } - Dim array_size() const { - return m_array_size; + /// The global offset (origin) of the chunk at the given grid index. + KMM_HOST_DEVICE + point_type chunk_offset(point_type grid_index) const noexcept { + point_type offset; + + for (size_t i = 0; is_less(i, N); i++) { + offset[i] = grid_index[i] * m_chunk_shape[i]; + } + + return offset; } - protected: - Dim m_array_size = Dim::zero(); - Dim m_chunk_size = Dim::zero(); - std::array m_chunks_count; - std::vector m_memories; -}; + /// Row-major linearization of a grid index (the last axis varies fastest). + KMM_HOST_DEVICE + size_t linear_index(point_type grid_index) const noexcept { + size_t linear = 0; + + for (size_t i = 0; is_less(i, N); i++) { + linear = + linear * static_cast(m_grid_shape[i]) + static_cast(grid_index[i]); + } -template -Distribution map_domain_to_distribution( - Dim array_size, - const Domain& domain, - M mapper, - bool allow_duplicates = false -) { - std::vector> chunks; - - for (const auto& chunk : domain.chunks) { - Bounds bounds = mapper(chunk, Bounds(array_size)); - - chunks.push_back(ArrayChunk { - .owner_id = chunk.owner_id.as_memory(), - .offset = bounds.begin(), - .size = bounds.size() - }); + return linear; } - return Distribution::from_chunks(array_size, chunks, allow_duplicates); -} + /// The range of grid indices (per axis, half-open `[begin, end)`) of the chunks that overlap + /// `region`. An axis with no overlap results in an empty range along that axis, which makes + /// the returned `Bounds` empty as a whole (see `Bounds::is_empty`). + KMM_HOST_DEVICE + Bounds chunk_range(const Bounds& region) const noexcept { + Bounds result; + + for (size_t i = 0; is_less(i, N); i++) { + auto lo = + region.begin(i) > static_cast(0) ? region.begin(i) : static_cast(0); + auto hi = region.end(i) < m_total_shape[i] ? region.end(i) : m_total_shape[i]; + + if (m_chunk_shape[i] > static_cast(0) && lo < hi) { + auto begin_chunk = lo / m_chunk_shape[i]; + auto end_chunk = + (hi - static_cast(1)) / m_chunk_shape[i] + static_cast(1); + + if (end_chunk > m_grid_shape[i]) { + end_chunk = m_grid_shape[i]; + } + + result[i] = Range {begin_chunk, end_chunk}; + } + } + + return result; + } + + /// Inverse of `linear_index`: recovers the grid index of the chunk with the given row-major + /// linear index. + KMM_HOST_DEVICE + point_type unravel(size_t linear) const noexcept { + point_type grid_index; + + for (size_t i = N; i-- > 0;) { + auto extent = static_cast(m_grid_shape[i]); + grid_index[i] = static_cast(extent > 0 ? linear % extent : 0); + linear = extent > 0 ? linear / extent : 0; + } + + return grid_index; + } + + private: + shape_type m_total_shape {}; + shape_type m_chunk_shape {}; + shape_type m_grid_shape {}; +}; -#define KMM_INSTANTIATE_ARRAY_IMPL(NAME) \ - template class NAME<0>; /* NOLINT */ \ - template class NAME<1>; /* NOLINT */ \ - template class NAME<2>; /* NOLINT */ \ - template class NAME<3>; /* NOLINT */ \ - template class NAME<4>; /* NOLINT */ \ - template class NAME<5>; /* NOLINT */ \ - template class NAME<6>; /* NOLINT */ +/// @} -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/core/domain.hpp b/include/kmm/core/domain.hpp deleted file mode 100644 index c850eb22..00000000 --- a/include/kmm/core/domain.hpp +++ /dev/null @@ -1,122 +0,0 @@ -#pragma once - -#include -#include - -#include "kmm/core/identifiers.hpp" -#include "kmm/core/resource.hpp" -#include "kmm/core/system_info.hpp" -#include "kmm/utils/geometry.hpp" - -namespace kmm { - -/** - * Constant for the number of dimensions in the work space. - */ -static constexpr size_t DOMAIN_DIMS = 3; - -/** - * Type alias for the index type used in the work space. - */ -class DomainPoint: public Point { - public: - DomainPoint(Point<0> p) : Point(p) {} - DomainPoint(Point<1> p) : Point(p) {} - DomainPoint(Point<2> p) : Point(p) {} - DomainPoint(Point<3> p) : Point(p) {} - - DomainPoint( - default_index_type x = 0, // - default_index_type y = 0, - default_index_type z = 0 - ) : - Point(x, y, z) {} -}; - -/** - * Type alias for the size of the work space. - */ -class DomainDim: public Dim { - public: - DomainDim(Dim<0> p) : Dim(p) {} - DomainDim(Dim<1> p) : Dim(p) {} - DomainDim(Dim<2> p) : Dim(p) {} - DomainDim(Dim<3> p) : Dim(p) {} - - DomainDim( - default_index_type x = 1, // - default_index_type y = 1, - default_index_type z = 1 - ) : - Dim(x, y, z) {} -}; - -/** - * Type alias for the size of the work space. - */ -class DomainBounds: public Bounds { - public: - DomainBounds(Bounds<0> p) : Bounds(p) {} - DomainBounds(Bounds<1> p) : Bounds(p) {} - DomainBounds(Bounds<2> p) : Bounds(p) {} - DomainBounds(Bounds<3> p) : Bounds(p) {} - - DomainBounds(Dim<0> p) : Bounds(p) {} - DomainBounds(Dim<1> p) : Bounds(p) {} - DomainBounds(Dim<2> p) : Bounds(p) {} - DomainBounds(Dim<3> p) : Bounds(p) {} - - DomainBounds( - Range x = 1, - Range y = 1, - Range z = 1 - ) : - Bounds(x, y, z) {} -}; - -struct DomainChunk { - ResourceId owner_id; - DomainPoint offset; - DomainDim size; -}; - -struct Domain { - std::vector chunks; -}; - -template -struct IntoDomain { - static Domain call(P partition, const SystemInfo& info, ExecutionSpace space) { - return partition; - } -}; - -struct TileDomain { - TileDomain(DomainDim domain_size, DomainDim tile_size) : - m_domain_size(domain_size), - m_tile_size(tile_size) {} - - Domain operator()(const SystemInfo& info, ExecutionSpace space) const; - - private: - DomainBounds m_domain_size; - DomainDim m_tile_size; -}; - -template<> -struct IntoDomain { - static Domain call(TileDomain partition, const SystemInfo& info, ExecutionSpace space) { - return partition(info, space); - } -}; - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/core/domain_traits.hpp b/include/kmm/core/domain_traits.hpp new file mode 100644 index 00000000..2b66177e --- /dev/null +++ b/include/kmm/core/domain_traits.hpp @@ -0,0 +1,12 @@ +#pragma once + +namespace kmm::detail { + +/// Customization point mapping a domain type (e.g. `Shape`, `Bounds`, `FShape`) to the interface +/// `Layout` needs from it: rank, index type, per-axis bounds/extent, and how the domain itself +/// changes under axis operations (`slice_axis`, `drop_axis`, `permute_axes`, `insert_axis`). +/// Specialized alongside each domain type, in that type's own header. +template +struct domain_traits; + +} // namespace kmm::detail diff --git a/include/kmm/core/fast_divisor.hpp b/include/kmm/core/fast_divisor.hpp new file mode 100644 index 00000000..5868996f --- /dev/null +++ b/include/kmm/core/fast_divisor.hpp @@ -0,0 +1,321 @@ +#pragma once + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/core/vec.hpp" + +#if !KMM_IS_RTC + #include +#endif + +namespace kmm { + +/// \addtogroup utility +/// @{ + +/// A precomputed divisor that replaces integer division by a cheaper multiply-and-shift. +/// +/// Uses the "round-up" magic-number scheme of T. Granlund and P. L. Montgomery, +/// "Division by Invariant Integers using Multiplication", ACM SIGPLAN PLDI 1994, +/// pp. 61-72 (https://doi.org/10.1145/178243.178249): the quotient is +/// `(numerator + mulhi(numerator, M)) >> shift`, where `M` is the low 32 (resp. 64) +/// bits of `ceil(2^(N+shift) / divisor)` and the always-set bit `N` is folded back in +/// as `+ numerator`. +template +class FastDivisor; + +template<> +class FastDivisor { + public: + /// The largest numerator that `divide` and `modulo` accept. + static constexpr uint32_t max_numerator = (uint32_t(1) << 31) - 1; + + /// Construct a divisor equal to one (the identity). + constexpr FastDivisor() = default; + + /// Precompute the magic constants for dividing by `divisor`. + /// + /// Throws `std::runtime_error` (or panics on device) when `divisor` is zero. + KMM_HOST_DEVICE + explicit constexpr FastDivisor(uint32_t divisor) : m_divisor(divisor) { + if (divisor == 0) { +#if KMM_IS_DEVICE + KMM_PANIC("FastDivisor: divisor must be non-zero"); +#elif !KMM_IS_RTC + throw std::runtime_error("FastDivisor: divisor must be non-zero"); +#endif + } + + // shift = ceil(log2(divisor)) + uint32_t shift = 0; + while ((uint64_t(1) << shift) < divisor) { + shift++; + } + + if (shift >= 32) { + // The divisor exceeds every admissible numerator, so the quotient is always 0. + m_multiplier = 0; + m_shift = 31; + } else { + // The magic number M = ceil(2^(32 + shift) / divisor) needs 33 bits; bit 32 is + // always set, so only the low 32 bits are stored and `divide` adds the numerator + // back to account for the implicit high bit. For a power of two M == 2^32 + // exactly, hence m_multiplier == 0 and `divide` reduces to a plain shift. + const uint64_t t = (uint64_t(1) << shift) - divisor; // 2^shift - divisor + m_multiplier = static_cast(((t << 32) + divisor - 1) / divisor); + m_shift = shift; + } + } + + /// Return `numerator / divisor`. Requires `numerator <= max_numerator`. + KMM_HOST_DEVICE + constexpr uint32_t divide(uint32_t numerator) const { + // KMM_ASSERT(numerator <= max_numerator); +#if KMM_IS_DEVICE + // `mulhi(a, b)` == high 32 bits of the 32x32 -> 64 bit product. + const uint32_t high = __umulhi(numerator, m_multiplier); +#else + const uint32_t high = static_cast((uint64_t(numerator) * m_multiplier) >> 32); +#endif + return (numerator + high) >> m_shift; + } + + /// Return `numerator % divisor`. Requires `numerator <= max_numerator`. + KMM_HOST_DEVICE + constexpr uint32_t modulo(uint32_t numerator) const { + return numerator - divide(numerator) * m_divisor; + } + + /// Return the divisor. + KMM_HOST_DEVICE + constexpr uint32_t get() const { + return m_divisor; + } + + private: + uint32_t m_divisor = 1; + uint32_t m_multiplier = 0; // low bits of the magic number; 0 for a power of two + uint32_t m_shift = 0; +}; + +template<> +class FastDivisor { + public: + /// The largest numerator that `divide` and `modulo` accept. + static constexpr uint64_t max_numerator = (uint64_t(1) << 63) - 1; + + /// Construct a divisor equal to one (the identity). + constexpr FastDivisor() = default; + + /// Precompute the magic constants for dividing by `divisor`. + /// + /// Throws `std::runtime_error` (or panics on device) when `divisor` is zero. + KMM_HOST_DEVICE + explicit constexpr FastDivisor(uint64_t divisor) : m_divisor(divisor) { + if (divisor == 0) { +#if KMM_IS_DEVICE + KMM_PANIC("FastDivisor: divisor must be non-zero"); +#elif !KMM_IS_RTC + throw std::runtime_error("FastDivisor: divisor must be non-zero"); +#endif + } + + // shift = ceil(log2(divisor)) + uint32_t shift = 0; + while (((unsigned __int128)(1) << shift) < divisor) { + shift++; + } + + if (shift >= 64) { + // The divisor exceeds every numerator, so division always returns zero. + m_multiplier = 0; + m_shift = 63; + } else { + // The magic number M = ceil(2^(64 + shift) / divisor) needs 65 bits; bit 64 is + // always set, so only the low 64 bits are stored and `divide` adds the numerator + // back to account for the implicit high bit. For a power of two M == 2^64 + // exactly, hence m_multiplier == 0 and `divide` reduces to a plain shift. + const unsigned __int128 t = ((unsigned __int128)(1) << shift) - divisor; + m_multiplier = static_cast(((t << 64) + divisor - 1) / divisor); + m_shift = shift; + } + } + + /// Return `numerator / divisor`. Requires `numerator <= max_numerator`. + KMM_HOST_DEVICE + constexpr uint64_t divide(uint64_t numerator) const { + // KMM_ASSERT(numerator <= max_numerator); +#if KMM_IS_DEVICE + // `mulhi(a, b)` == high 64 bits of the 64x64 -> 128 bit product. + const uint64_t high = __umul64hi(numerator, m_multiplier); +#else + const uint64_t high = + static_cast(((unsigned __int128)(numerator)*m_multiplier) >> 64); +#endif + return (numerator + high) >> m_shift; + } + + /// Return `numerator % divisor`. Requires `numerator <= max_numerator`. + KMM_HOST_DEVICE + constexpr uint64_t modulo(uint64_t numerator) const { + return numerator - divide(numerator) * m_divisor; + } + + /// Return the divisor. + KMM_HOST_DEVICE + constexpr uint64_t get() const { + return m_divisor; + } + + private: + uint64_t m_divisor = 1; + uint64_t m_multiplier = 0; // low bits of the magic number; 0 for a power of two + uint32_t m_shift = 0; +}; + +template +KMM_HOST_DEVICE constexpr T operator/(T numerator, const FastDivisor& divisor) { + return divisor.divide(numerator); +} + +template +KMM_HOST_DEVICE constexpr T operator%(T numerator, const FastDivisor& divisor) { + return divisor.modulo(numerator); +} + +/// Converts between a flat index and its `N`-dimensional coordinates and back. +/// +/// The conversion is `flat_index == coord[0] + extent[0] * (coord[1] + extent[1] * (... + extent[N-2] * coord[N-1]))`. +/// This means extent[0] is the most-contiguous index. +template +class IndexMapper { + public: + /// The largest `volume()` that `unravel` supports. + static constexpr T max_volume = FastDivisor::max_numerator + T(1); + + IndexMapper() = default; + + KMM_HOST_DEVICE + explicit IndexMapper(const T* extents) { + for (size_t i = 0; i < N - 1; i++) { + m_extents[i] = FastDivisor(extents[i]); + } + + m_volume = checked_product(extents, extents + N); + } + + KMM_HOST_DEVICE + explicit IndexMapper(const Vec& extents) : IndexMapper(&extents[0]) {} + + KMM_HOST_DEVICE + T volume() const noexcept { + return m_volume; + } + + /// Decompose `index` into `result`. Returns `true` if `index` lies within the space. + /// + /// `result` is fully overwritten, so the caller need not initialize it. + KMM_HOST_DEVICE + bool unravel(T linear_index, Vec& result) const { + if (linear_index >= m_volume) { + return false; + } + + for (size_t i = 0; i < N - 1; i++) { + const T quo = linear_index / m_extents[i]; + result[i] = linear_index - quo * m_extents[i].get(); + linear_index = quo; + } + + result[N - 1] = linear_index; + return true; + } + + /// Combine the coordinates in `coord` back into a flat index. + KMM_HOST_DEVICE + T ravel(const Vec& coord) const { + T linear_index = coord[N - 1]; + for (size_t i = N - 1; i-- > 0;) { + linear_index = linear_index * m_extents[i].get() + coord[i]; + } + return linear_index; + } + + private: + FastDivisor m_extents[N - 1]; // we do not need store the last extent + T m_volume = 1; +}; + +template +class IndexMapper { + public: + /// A 1-D mapping performs no division, so the volume is limited only by `T`. + static constexpr T max_volume = ~T(0); + + IndexMapper() = default; + + KMM_HOST_DEVICE + explicit IndexMapper(const T* extents) { + m_extent = extents[0]; + } + + KMM_HOST_DEVICE + explicit IndexMapper(const Vec& extents) : IndexMapper(&extents.x) {} + + KMM_HOST_DEVICE + T volume() const noexcept { + return m_extent; + } + + /// Decompose `index` into `result`. Returns `true` if `index` lies within the space. + /// + /// `result` is fully overwritten, so the caller need not initialize it. + KMM_HOST_DEVICE + bool unravel(T linear_index, Vec& result) const { + result[0] = linear_index; + return linear_index < m_extent; + } + + /// Combine the coordinates in `coord` back into a flat index. + KMM_HOST_DEVICE + T ravel(const Vec& coord) const { + return coord[0]; + } + + private: + T m_extent = 1; +}; + +template +class IndexMapper { + public: + /// A 0-D mapping performs no division, so the volume is limited only by `T`. + static constexpr T max_volume = ~T(0); + + IndexMapper() = default; + + KMM_HOST_DEVICE + explicit IndexMapper(const T*) {} + + KMM_HOST_DEVICE + explicit IndexMapper(const Vec& extents) {} + + KMM_HOST_DEVICE + T volume() const noexcept { + return 0; + } + + KMM_HOST_DEVICE + bool unravel(T index, Vec&) const { + return index == 0; + } + + KMM_HOST_DEVICE + T ravel(const Vec&) const { + return 0; + } +}; + +/// @} + +} // namespace kmm diff --git a/include/kmm/core/fshape.hpp b/include/kmm/core/fshape.hpp new file mode 100644 index 00000000..af784507 --- /dev/null +++ b/include/kmm/core/fshape.hpp @@ -0,0 +1,318 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/domain_traits.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/range.hpp" +#include "kmm/core/type_utils.hpp" +#include "kmm/core/vec.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// The extent (size) of an N-dimensional domain along each axis, backed by a `Vec`, using +/// Fortran-style 1-based indexing: axis `i` covers the valid indices `[1, extent]` instead of the +/// C-style `[0, extent)` used by `Shape`. +/// +/// `FShape` is structurally identical to `Shape` (same storage and extent semantics, i.e. +/// `extent(i)` still just means "how many elements along axis `i`"); the indexing convention only +/// takes effect through the `detail::domain_traits>` specialization used by `Layout`. +template +class FShape: public Vec { + public: + using storage_type = Vec; + + KMM_HOST_DEVICE + explicit constexpr FShape(const storage_type& storage) : storage_type(storage) {} + + /// Create an empty shape `(0, 0, ...)` + KMM_HOST_DEVICE + constexpr FShape() : storage_type(::kmm::fill(static_cast(0))) {} + + constexpr FShape(const FShape&) = default; + constexpr FShape(FShape&&) noexcept = default; + FShape& operator=(const FShape&) = default; + FShape& operator=(FShape&&) noexcept = default; + + /// Create a shape `(first, args...)` + template> + KMM_HOST_DEVICE FShape(T first, Ts&&... args) : storage_type {first, args...} {} + + /// Create a shape from another shape. Throws on overflow. + template + KMM_HOST_DEVICE constexpr FShape(const FShape& that) { + if (!that.template is_convertible_to()) { + throw_overflow_exception(); + } + + *this = FShape::from(that); + } + + /// Create a shape from another shape. Does not throw on overflow. + template + KMM_HOST_DEVICE static constexpr FShape from(const Vec& that) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = is_less(i, M) ? static_cast(that[i]) : static_cast(1); + } + + return FShape(result); + } + + /// Create a shape `(value, value, value, ...)`. + KMM_HOST_DEVICE + static constexpr FShape fill(T value) { + return FShape(::kmm::fill(value)); + } + + /// Create a shape `(1, 1, 1, ...)`. + KMM_HOST_DEVICE + static constexpr FShape one() { + return fill(static_cast(1)); + } + + /// Create a shape `(0, 0, 0, ...)`. + KMM_HOST_DEVICE + static constexpr FShape zero() { + return fill(static_cast(0)); + } + + /// Returns coordinate i, or default_value if the axis is out of range. + KMM_HOST_DEVICE + T get_or_default(size_t i, T default_value = static_cast(1)) const { + if constexpr (N > 0) { + if (KMM_LIKELY(is_less(i, N))) { + return (*this)[i]; + } + } + + return default_value; + } + + /// Checks whether this shape can be converted to `FShape`. + template + KMM_HOST_DEVICE bool is_convertible_to() const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + if (i < M) { + result &= is_convertible((*this)[i]); + } else { + result &= is_equal((*this)[i], static_cast(1)); + } + } + + return result; + } + + /// Check if this shape is empty. + /// + /// A shape is empty if each axis is less than or equal to zero. + KMM_HOST_DEVICE + bool is_empty() const { + bool result = false; + + for (size_t i = 0; is_less(i, N); i++) { + result |= !(static_cast(0) < (*this)[i]); + } + + return result; + } + + /// Returns the product of the extents of this shape. + /// + /// If one of the axis is negative, the returned value is zero. + KMM_HOST_DEVICE + T volume() const { + T result = static_cast(1); + + if constexpr (N >= 1) { + result = (*this)[0]; + + for (size_t i = 1; is_less(i, N); i++) { + result *= (*this)[i]; + } + } + + return is_empty() ? static_cast(0) : result; + } + + /// Check if a point falls within this shape, using 1-based indices. + /// + /// This means that for each axis `p[i] >= 1` and `p[i] <= shape[i]`. + template + KMM_HOST_DEVICE bool contains(const Point& p) const { + bool result = true; + + for (size_t i = 0; is_less(i, N) && is_less(i, M); i++) { + result &= !is_less(p[i], static_cast(1)) && !is_less((*this)[i], p[i]); + } + + if constexpr (N < M) { + for (size_t i = N; is_less(i, M); i++) { + result &= is_equal(p[i], static_cast(1)); + } + } + + if constexpr (N > M) { + for (size_t i = M; is_less(i, N); i++) { + result &= is_less(static_cast(0), (*this)[i]); + } + } + + return result; + } +}; + +template +FShape(Ts&&...) -> FShape; + +/// Constructs an FShape from the given per-axis extents. +template +KMM_HOST_DEVICE FShape fshape(const Ts&... values) { + return FShape { + Vec {static_cast(values)...} + }; +} + +template +KMM_HOST_DEVICE FShape concat(const FShape& lhs, const FShape& rhs) { + return FShape { + concat(static_cast&>(lhs), static_cast&>(rhs)) + }; +} + +template +KMM_HOST_DEVICE bool operator==(const FShape& lhs, const FShape& rhs) { + bool result = true; + + for (size_t i = 0; is_less(i, N) || is_less(i, M); i++) { + result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); + } + + return result; +} + +template +KMM_HOST_DEVICE bool operator!=(const FShape& lhs, const FShape& rhs) { + return !(lhs == rhs); +} + +/// @} + +namespace detail { + +// Same extent semantics as `domain_traits>` -- only `bounds()` differs, since +// `FShape` is Fortran-style 1-based: valid indices along an axis are `[1, extent]` rather than +// `[0, extent)`. `extent`/`slice_axis`/`drop_axis`/`permute_axes`/`insert_axis` all operate on +// extents (counts), which are indexing-convention-agnostic, so their logic is identical to +// `Shape`'s, just retargeted to produce `FShape` instead of `Shape`. +template +struct domain_traits> { + static constexpr size_t rank = N; + using index_type = IndexT; + using domain_type = FShape; + + KMM_HOST_DEVICE + static constexpr Range bounds(const domain_type& domain, size_t axis) { + return {static_cast(1), domain[axis] + static_cast(1)}; + } + + KMM_HOST_DEVICE + static constexpr index_type extent(const domain_type& domain, size_t axis) { + return domain[axis]; + } + + template + using slice_axis_type = FShape; + + template + KMM_HOST_DEVICE static constexpr slice_axis_type slice_axis( + domain_type domain, + index_type begin, + index_type end + ) { + domain[Axis] = end - begin; + return domain; + } + template + using drop_axis_type = FShape; + + template + KMM_HOST_DEVICE static constexpr drop_axis_type drop_axis(const domain_type& domain) { + return permute_axes(domain, drop_index_sequence()); + } + + /// Runtime-axis counterpart of `drop_axis`. The result type only depends on `N` (not on + /// which axis is dropped), so `axis` does not need to be known at compile time here. + KMM_HOST_DEVICE + static constexpr drop_axis_type<0> drop_axis(const domain_type& domain, size_t axis) { + KMM_ASSERT(axis < N); + drop_axis_type<0> result; + size_t j = 0; + + for (size_t i = 0; is_less(i, N); i++) { + if (i != axis) { + result[j] = domain[i]; + j++; + } + } + + return result; + } + + template + using permute_axes_type = FShape; + + template + KMM_HOST_DEVICE static constexpr permute_axes_type permute_axes( + const domain_type& domain, + IndexSequence + ) { + return {domain[Is]...}; + } + + template + using insert_axis_type = FShape; + + template + KMM_HOST_DEVICE static constexpr insert_axis_type insert_axis( + const domain_type& domain, + index_type extent + ) { + insert_axis_type result; + + for (size_t i = 0; is_less(i, Axis); i++) { + result[i] = domain[i]; + } + + result[Axis] = extent; + + for (size_t i = Axis; is_less(i, N); i++) { + result[i + 1] = domain[i]; + } + + return result; + } +}; + +} // namespace detail + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash>: std::hash> {}; +#endif diff --git a/include/kmm/core/identifiers.hpp b/include/kmm/core/identifiers.hpp deleted file mode 100644 index 5f73e7a7..00000000 --- a/include/kmm/core/identifiers.hpp +++ /dev/null @@ -1,368 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "fmt/ostream.h" - -#include "kmm/utils/checked_math.hpp" -#include "kmm/utils/macros.hpp" -#include "kmm/utils/panic.hpp" -#include "kmm/utils/small_vector.hpp" - -#define KMM_IMPL_COMPARISON_OPS(T) \ - KMM_INLINE constexpr bool operator!=(const T& that) const { \ - return !(*this == that); \ - } \ - KMM_INLINE constexpr bool operator<=(const T& that) const { \ - return !(*this > that); \ - } \ - KMM_INLINE constexpr bool operator>(const T& that) const { \ - return that < *this; \ - } \ - KMM_INLINE constexpr bool operator>=(const T& that) const { \ - return that <= *this; \ - } - -namespace kmm { - -using index_t = int; - -struct NodeId { - explicit constexpr NodeId(uint8_t v) : m_value(v) {} - explicit constexpr NodeId(size_t v) : m_value(checked_cast(v)) {} - - KMM_INLINE constexpr uint8_t get() const { - return m_value; - } - - KMM_INLINE operator uint8_t() const { - return get(); - } - - KMM_INLINE constexpr bool operator==(const NodeId& that) const { - return m_value == that.m_value; - } - - KMM_INLINE constexpr bool operator<(const NodeId& that) const { - return m_value < that.m_value; - } - - KMM_IMPL_COMPARISON_OPS(NodeId) - - friend std::ostream& operator<<(std::ostream&, const NodeId&); - - private: - uint8_t m_value; -}; - -// Maximum of 8 devices per node -static constexpr size_t MAX_DEVICES = 8; - -struct DeviceId { - template - explicit constexpr DeviceId(T v) : m_value(static_cast(v)) { - if (!in_range(v, MAX_DEVICES)) { - throw std::runtime_error("device index out of range"); - } - } - - KMM_INLINE constexpr uint8_t get() const { - if (m_value >= MAX_DEVICES) { - __builtin_unreachable(); - } - - return m_value; - } - - KMM_INLINE operator uint8_t() const { - return get(); - } - - KMM_INLINE constexpr bool operator==(const DeviceId& that) const { - return m_value == that.m_value; - } - - KMM_INLINE constexpr bool operator<(const DeviceId& that) const { - return m_value < that.m_value; - } - - KMM_IMPL_COMPARISON_OPS(DeviceId) - - friend std::ostream& operator<<(std::ostream&, const DeviceId&); - - private: - uint8_t m_value; -}; - -struct DeviceStreamSet { - static constexpr size_t MAX_SIZE = 64; - using value_type = uint64_t; - - constexpr DeviceStreamSet(size_t index) { - if (index < MAX_SIZE) { - this->m_value |= value_type(1) << value_type(index); - } - } - - constexpr DeviceStreamSet() = default; - - constexpr DeviceStreamSet(std::initializer_list indices) { - for (auto index : indices) { - this->m_value |= DeviceStreamSet(index).m_value; - } - } - - static constexpr DeviceStreamSet range(size_t begin, size_t end) { - auto result = DeviceStreamSet {}; - - for (size_t index = begin; index < end; index++) { - result.m_value |= DeviceStreamSet(index).m_value; - } - - return result; - } - - static constexpr DeviceStreamSet all() { - return range(0, MAX_SIZE); - } - - constexpr bool is_empty() const { - return m_value == 0; - } - - constexpr bool contains(DeviceStreamSet that) const { - return (this->m_value & that.m_value) == that.m_value; - } - - constexpr bool contains(size_t index) const { - return contains(DeviceStreamSet {index}); - } - - constexpr DeviceStreamSet& operator&=(const DeviceStreamSet& that) { - this->m_value &= that.m_value; - return *this; - } - - constexpr DeviceStreamSet operator&(const DeviceStreamSet& that) const { - return DeviceStreamSet {*this} &= that; - } - - constexpr bool operator==(const DeviceStreamSet& that) const { - return m_value == that.m_value; - } - - constexpr bool operator!=(const DeviceStreamSet& that) const { - return !(*this == that); - } - - private: - value_type m_value = 0; -}; - -struct MemoryId { - public: - enum struct Type : uint8_t { Host, Device }; - - KMM_INLINE constexpr MemoryId(Type kind, uint8_t v) : m_type(kind), m_value(v) {} - - KMM_INLINE constexpr MemoryId(DeviceId device) : MemoryId(Type::Device, device.get()) {} - - KMM_INLINE static constexpr MemoryId host(DeviceId device_affinity = DeviceId(0)) { - return MemoryId {Type::Host, device_affinity.get()}; - } - - KMM_INLINE constexpr bool is_host() const noexcept { - return m_type == Type::Host; - } - - KMM_INLINE constexpr bool is_device() const noexcept { - return m_type == Type::Device; - } - - KMM_INLINE constexpr DeviceId as_device() const { - KMM_ASSERT(is_device()); - return DeviceId(m_value); - } - - KMM_INLINE constexpr DeviceId device_affinity() const { - return DeviceId(m_value % MAX_DEVICES); - } - - KMM_INLINE constexpr bool operator==(const MemoryId& that) const { - if (m_type == Type::Device && m_type == Type::Device) { - return m_value == that.m_value; - } else { - return m_type == that.m_type; - } - } - - KMM_INLINE constexpr bool operator<(const MemoryId& that) const { - if (m_type == Type::Device && m_type == Type::Device) { - return m_value < that.m_value; - } else { - return m_type < that.m_type; - } - } - - KMM_IMPL_COMPARISON_OPS(MemoryId) - - friend std::ostream& operator<<(std::ostream&, const MemoryId&); - - private: - Type m_type; - uint8_t m_value; -}; - -class ResourceId { - public: - enum struct Type : uint8_t { Host, Device }; - - KMM_INLINE constexpr ResourceId() : m_type(Type::Host) {} - - KMM_INLINE constexpr ResourceId( - DeviceId device, - DeviceStreamSet stream_affinity = DeviceStreamSet::all() - ) : - m_type(Type::Device), - m_device(device), - m_stream_affinity(stream_affinity) {} - - KMM_INLINE static constexpr ResourceId host(DeviceId affinity = DeviceId(0)) { - ResourceId result; - result.m_device = affinity; - return result; - } - - KMM_INLINE constexpr bool is_host() const { - return m_type == Type::Host; - } - - KMM_INLINE constexpr bool is_device() const { - return m_type == Type::Device; - } - - KMM_INLINE constexpr DeviceId as_device() const { - KMM_ASSERT(is_device()); - return m_device; - } - - KMM_INLINE constexpr DeviceId device_affinity() const { - return m_device; - } - - KMM_INLINE constexpr DeviceStreamSet stream_affinity() const { - if (m_type == Type::Device && m_stream_affinity != DeviceStreamSet {}) { - return m_stream_affinity; - } else { - return DeviceStreamSet::all(); - } - } - - KMM_INLINE constexpr MemoryId as_memory() const { - return is_host() ? MemoryId::host(m_device) : MemoryId(m_device); - } - - KMM_INLINE constexpr bool contains(const ResourceId& that) const { - if (m_type == Type::Host && that.m_type == Type::Host) { - return true; - } else if (m_type == Type::Device && that.m_type == Type::Device) { - return m_device == that.m_device && m_stream_affinity.contains(that.m_stream_affinity); - } else { - return false; - } - } - - friend std::ostream& operator<<(std::ostream&, const ResourceId&); - - private: - Type m_type = Type::Host; - DeviceId m_device = DeviceId(0); - DeviceStreamSet m_stream_affinity; // only for Type::Device -}; - -struct BufferId { - KMM_INLINE explicit constexpr BufferId(uint64_t v = ~uint64_t(0)) : m_value(v) {} - - KMM_INLINE constexpr uint64_t get() const { - return m_value; - } - - KMM_INLINE operator uint64_t() const { - return get(); - } - - KMM_INLINE constexpr bool operator==(const BufferId& that) const { - return m_value == that.m_value; - } - - KMM_INLINE constexpr bool operator<(const BufferId& that) const { - return m_value < that.m_value; - } - - KMM_IMPL_COMPARISON_OPS(BufferId) - - friend std::ostream& operator<<(std::ostream&, const BufferId&); - - private: - uint64_t m_value; -}; - -struct EventId { - KMM_INLINE constexpr EventId() = default; - - KMM_INLINE explicit constexpr EventId(uint64_t v) : m_value(v) {} - - KMM_INLINE constexpr uint64_t get() const { - return m_value; - } - - KMM_INLINE operator uint64_t() const { - return get(); - } - - KMM_INLINE constexpr bool operator==(const EventId& that) const { - return m_value == that.m_value; - } - - KMM_INLINE constexpr bool operator<(const EventId& that) const { - return m_value < that.m_value; - } - - KMM_IMPL_COMPARISON_OPS(EventId) - - friend std::ostream& operator<<(std::ostream&, const EventId&); - - private: - uint64_t m_value = 0; -}; - -using EventList = small_vector; -std::ostream& operator<<(std::ostream&, const EventList&); - -} // namespace kmm - -template<> -struct std::hash: std::hash {}; -template<> -struct std::hash: std::hash {}; -template<> -struct std::hash: std::hash {}; -template<> -struct std::hash: std::hash {}; - -template<> -struct fmt::formatter: fmt::formatter {}; -template<> -struct fmt::formatter: fmt::formatter {}; -template<> -struct fmt::formatter: fmt::formatter {}; -template<> -struct fmt::formatter: fmt::formatter {}; -template<> -struct fmt::formatter: fmt::ostream_formatter {}; -template<> -struct fmt::formatter: fmt::ostream_formatter {}; -template<> -struct fmt::formatter: fmt::ostream_formatter {}; diff --git a/include/kmm/core/integer_fun.hpp b/include/kmm/core/integer_fun.hpp new file mode 100644 index 00000000..71f7d921 --- /dev/null +++ b/include/kmm/core/integer_fun.hpp @@ -0,0 +1,103 @@ +#pragma once + +#include "kmm/core/checked_math.hpp" +#include "kmm/core/macros.hpp" + +namespace kmm { + +/// \addtogroup utility +/// @{ + +/// Divide `num` by `denom` and round the result down (towards negative infinity). +template +KMM_HOST_DEVICE constexpr T div_floor(T a, T b) { + const T zero = static_cast(0); + T quotient = a / b; + + if constexpr (detail::numeric_type_traits::is_signed) { + // Adjust the quotient if a and b have different signs + if (a % b != zero && ((a >= zero) ^ (b >= zero))) { + quotient -= 1; + } + } + + return quotient; +} + +/// Divide `num` by `denom` and round the result up (towards positive infinity). +template +KMM_HOST_DEVICE constexpr T div_ceil(T a, T b) { + const T zero = static_cast(0); + T quotient = a / b; + + if constexpr (detail::numeric_type_traits::is_signed) { + // Adjust the quotient if both a and b have the same sign + if (a % b != zero && !((a >= zero) ^ (b >= zero))) { + quotient += 1; + } + } else if (a % b != zero) { + quotient += 1; + } + + return quotient; +} + +/// Return the absolute value of `input` as the corresponding unsigned integer type. +template::unsigned_type> +KMM_HOST_DEVICE constexpr U unsigned_abs(T input) { + U magnitude = static_cast(input); + return detail::is_negative(input) ? static_cast(U(0) - magnitude) : magnitude; +} + +/// Round `input` to the next multiple of `multiple`. +/// +/// This returns the smallest value greater than or equal to `input` that is divisible +/// by `multiple`. +template +KMM_HOST_DEVICE constexpr T round_up_to_multiple(T input, T multiple) { + using U = typename detail::numeric_type_traits::unsigned_type; + const T zero = static_cast(0); + + T remainder = input % multiple; + T delta = zero; + + if (remainder != zero) { + U magnitude = unsigned_abs(multiple) * !detail::is_negative(input); + delta = static_cast(static_cast(magnitude - static_cast(remainder))); + } + + return checked_add(input, delta); +} + +/// Returns `true` if `left` is exactly divisible by `right` (i.e. `left % right == 0`). +/// Returns `false` otherwise. This function accepts mixed inputs having different signedness. +template +KMM_HOST_DEVICE constexpr bool is_divisible(L left, R right) { + return right != R {0} && unsigned_abs(left) % unsigned_abs(right) == 0U; +} + +/// Returns `true` if `input` is a power of two (i.e. `1`, `2`, `4`, `8`, ...). +/// Returns `false` otherwise, including for `input <= 0`. +template +KMM_HOST_DEVICE constexpr bool is_power_of_two(T input) { + return input > static_cast(0) && (input & (input - static_cast(1))) == static_cast(0); +} + +/// Return the smallest integer that is a power of two and is not less than `input`. +template +KMM_HOST_DEVICE constexpr T round_up_to_power_of_two(T input) { + if (input <= static_cast(0)) { + return static_cast(1); + } + + input -= static_cast(1); + for (decltype(sizeof(T)) i = 1; i < sizeof(T) * 8; i *= 2) { + input |= (input >> i); + } + + return checked_add(input, static_cast(1)); +} + +/// @} + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/utils/key_value.hpp b/include/kmm/core/key_value.hpp similarity index 52% rename from include/kmm/utils/key_value.hpp rename to include/kmm/core/key_value.hpp index 80ae01dc..9d0492fd 100644 --- a/include/kmm/utils/key_value.hpp +++ b/include/kmm/core/key_value.hpp @@ -1,20 +1,23 @@ #pragma once -#include "kmm/utils/macros.hpp" +#include "kmm/core/macros.hpp" namespace kmm { -/** - * Pair of a key (type `int64_t`) and an associated value (type `T`). - * - * The main purpose is that `KeyValue` pairs implement the `>` operator, allowed them to be sorted. - * The pairs are ordered by value, with ties resolved by the key. - */ -template -struct alignas(alignof(long) >= alignof(V) ? 2 * alignof(long) : alignof(V)) KeyValue { +/// \addtogroup utility +/// @{ + +/// Pair of a value (type `T`) and an associated key (type `int64_t`). +/// +/// The main purpose is that `KeyValue` pairs implement the `<` operator, allowed them to be sorted. +/// The pairs are ordered by value, with ties resolved by the key. This is useful to sort a list +/// of `KeyValue` pairs by value and still have the associated key. +template +struct alignas(2 * alignof(long)) KeyValue { using key_type = long; - using value_type = V; + using value_type = ValueT; + KMM_HOST_DEVICE constexpr KeyValue() = default; KMM_HOST_DEVICE @@ -26,22 +29,17 @@ struct alignas(alignof(long) >= alignof(V) ? 2 * alignof(long) : alignof(V)) Key template KMM_HOST_DEVICE bool operator==(const KeyValue& a, const KeyValue& b) { - return (a.value == b.value) & (a.key == b.key); + return (a.key == b.key) && (a.value == b.value); } template KMM_HOST_DEVICE bool operator<(const KeyValue& a, const KeyValue& b) { - // The obvious expression would be: - // `(a.value < b.value` || (a.value == b.value && a.key < b.key)`. - // However, this would no work well with `NaN`, since `a.value == b.value` would then - // become false. Instead, we use `!(b.value < a.value)` the check if `a` and `b` have - // the same value, since this properly gives `true` for `NaN` values. - return (a.value < b.value) | ((!(b.value < a.value)) & (a.key < b.key)); + return (a.value < b.value) | ((a.value == b.value) & (a.key < b.key)); } template KMM_HOST_DEVICE bool operator<=(const KeyValue& a, const KeyValue& b) { - return (a.value < b.value) | ((!(b.value < a.value)) & (a.key <= b.key)); + return (a.value < b.value) | ((a.value == b.value) & (a.key <= b.key)); } template @@ -59,20 +57,21 @@ KMM_HOST_DEVICE bool operator>=(const KeyValue& a, const KeyValue& b) { return b <= a; } +/// @} + } // namespace kmm #if !KMM_IS_RTC - #include + #include #include "fmt/ostream.h" #include "kmm/utils/hash_utils.hpp" namespace kmm { - template -std::ostream& operator<<(std::ostream& stream, const KeyValue& p) { - return stream << "{key=" << p.key << ", value=" << p.value << "}"; +std::ostream& operator<<(std::ostream& stream, const KeyValue& kv) { + return stream << "{key=" << kv.key << ", value=" << kv.value << "}"; } } // namespace kmm @@ -81,8 +80,8 @@ struct fmt::formatter>: fmt::ostream_formatter {}; template struct std::hash> { - size_t operator()(const kmm::KeyValue& p) const { - return kmm::hash_fields(p.key, p.value); + size_t operator()(const kmm::KeyValue& kv) const { + return kmm::hash_fields(kv.key, kv.value); } }; #endif \ No newline at end of file diff --git a/include/kmm/core/layout.hpp b/include/kmm/core/layout.hpp new file mode 100644 index 00000000..50f44bd2 --- /dev/null +++ b/include/kmm/core/layout.hpp @@ -0,0 +1,897 @@ +#pragma once + +#include "kmm/core/bounds.hpp" +#include "kmm/core/const_value.hpp" +#include "kmm/core/fshape.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/range.hpp" +#include "kmm/core/shape.hpp" +#include "kmm/core/strides.hpp" +#include "kmm/core/type_utils.hpp" + +namespace kmm { + +struct all_t { + explicit all_t() = default; +}; +constexpr static all_t all = all_t(); + +struct new_axis_t { + explicit new_axis_t() = default; +}; +constexpr static new_axis_t new_axis = new_axis_t(); + +namespace detail { + +template +struct policy_traits; + +template< + typename PolicyT, + typename DomainT, + typename MappingT = typename policy_traits::mapping_type> +struct mapping_traits; + +template +struct mapping_traits> { + static constexpr size_t rank = sizeof...(StridesT); + using mapping_type = Strides; + using stride_type = default_stride_type; + + KMM_HOST_DEVICE + static constexpr stride_type stride(const mapping_type& mapping, size_t axis) { + return mapping[axis]; + } + + template + KMM_HOST_DEVICE static constexpr ptrdiff_t linearize_offset( + const mapping_type& mapping, + const Vec& index + ) { + return mapping.linearize_offset(index); + } + + template + struct unpack_sequence_helper; + + template + struct unpack_sequence_helper> { + using type = Strides...>; + + KMM_HOST_DEVICE + static constexpr type apply(const mapping_type& mapping) { + return type {mapping.get(ConstIndex())...}; + } + }; + + template + using drop_axis_type = typename unpack_sequence_helper>::type; + + template + KMM_HOST_DEVICE static constexpr drop_axis_type drop_axis(const mapping_type& mapping) { + static_assert(Axis < rank, "axis out of bounds"); + return unpack_sequence_helper>::apply(mapping); + } + + template + using permute_axes_type = Strides...>; + + template + KMM_HOST_DEVICE static constexpr permute_axes_type permute_axes( + const mapping_type& mapping, + IndexSequence + ) { + return permute_axes_type {mapping.get(ConstIndex())...}; + } + + template + struct insert_axis_helper; + + // The inserted axis always gets a compile-time-static stride of 0: it isn't backed by any + // real memory dimension, so every index along it must map to the same offset (broadcast). + template + struct insert_axis_helper, IndexSequence> { + using type = Strides< + typename mapping_type::template axis_stride_type..., + ConstValue, + typename mapping_type::template axis_stride_type...>; + + KMM_HOST_DEVICE + static constexpr type apply(const mapping_type& mapping) { + return type { + mapping.get(ConstIndex())..., + ConstValue {}, + mapping.get(ConstIndex())... + }; + } + }; + + template + using insert_axis_type = typename insert_axis_helper< + Axis, + make_index_sequence, + range_index_sequence_t>::type; + + template + KMM_HOST_DEVICE static constexpr insert_axis_type insert_axis( + const mapping_type& mapping + ) { + static_assert(Axis <= rank, "axis out of bounds"); + return insert_axis_helper< + Axis, + make_index_sequence, + range_index_sequence_t>::apply(mapping); + } +}; + +/// Customization point mapping a stride policy (e.g. `RowMajor`, `ColMajor`) and a domain type to +/// the mapping it produces. The default definition below just forwards to `PolicyT`'s own +/// `mapping_type` member alias and `apply` member function, so any policy following that +/// (intrusive) convention keeps working unchanged; specialize this trait instead when a policy +/// type can't (or shouldn't) define those members itself, e.g. a foreign type. +template +struct policy_traits { + using mapping_type = typename PolicyT::template mapping_type; + + KMM_HOST_DEVICE + static mapping_type apply(const PolicyT& policy, const DomainT& domain) { + return policy.apply(domain); + } +}; + +/// A `Strides<...>` used directly as a policy is its own mapping, independent of the domain: this +/// lets a `Layout` be constructed from an already-computed mapping (as if it were a policy) +/// instead of only from a policy that derives one from the domain. +template +struct policy_traits, DomainT> { + using mapping_type = Strides; + + KMM_HOST_DEVICE + static mapping_type apply(const Strides& policy, const DomainT&) { + return policy; + } +}; + +template +struct slice_axis_impl; + +template +struct slice_axis_impl { + // How many axes this token writes into the result at this position; slice_multi_impl uses + // this to compute where the next token's axis index lands (next_axis = Axis + produced_axes). + static constexpr size_t produced_axes = 1; + using type = LayoutT; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, all_t) { + return layout; + } +}; + +template +struct slice_axis_impl { + static constexpr size_t produced_axes = 0; + using type = typename LayoutT::template drop_axis_type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, int index) { + return layout.template drop_axis(static_cast(index)); + } +}; + +template +struct slice_axis_impl { + static constexpr size_t produced_axes = 0; + using type = typename LayoutT::template drop_axis_type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, long index) { + return layout.template drop_axis(static_cast(index)); + } +}; + +template +struct slice_axis_impl { + static constexpr size_t produced_axes = 0; + using type = typename LayoutT::template drop_axis_type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, size_t index) { + return layout.template drop_axis(static_cast(index)); + } +}; + +template +struct slice_axis_impl> { + static constexpr size_t produced_axes = 1; + using type = LayoutT; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, Range range) { + using index_type = typename LayoutT::index_type; + return layout.template slice_axis( + static_cast(range.start), + static_cast(range.stop) + ); + } +}; + +template +struct slice_axis_impl { + static constexpr size_t produced_axes = 1; + using type = typename LayoutT::template insert_axis_type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, new_axis_t) { + return layout.template insert_axis(); + } +}; + +template +struct slice_multi_impl; + +// Base case: every slice token has been consumed. +template +struct slice_multi_impl { + static_assert(Axis == LayoutT::rank, "invalid number of axis slices"); + using type = LayoutT; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout) { + return layout; + } +}; + +// General case: consume the head slice token (HeadSliceT), then recurse on the rest (TailSlicesT...). +template +struct slice_multi_impl { + // Bounds-checking against LayoutT::rank is left to slice_axis_impl::apply itself (via + // drop_axis/slice_axis/insert_axis's own static_asserts), since whether an existing axis is + // required at all depends on the token (e.g. new_axis_t needs none). + static_assert(Axis <= LayoutT::rank, "axis out of bounds"); + + using new_layout = typename slice_axis_impl::type; + static constexpr size_t next_axis = + Axis + slice_axis_impl::produced_axes; + using type = typename slice_multi_impl::type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, const SliceT& head, const Rest&... tail) { + return slice_multi_impl::apply( + slice_axis_impl::apply(layout, head), + tail... + ); + } +}; + +// Specialized for performance: `all_t` is a no-op on the layout, so just skip straight to the +// next axis without going through slice_axis_impl<..., all_t>::apply at all. +template +struct slice_multi_impl { + using type = typename slice_multi_impl::type; + + KMM_HOST_DEVICE + static type apply(const LayoutT& layout, all_t, const Rest&... tail) { + return slice_multi_impl::apply(layout, tail...); + } +}; + +} // namespace detail + +template +using domain_index_type = typename detail::domain_traits::index_type; + +template +static constexpr size_t domain_rank = detail::domain_traits::rank; + +namespace detail { + +// Calls make_strides(extent_0, extent_1, ...), reading each axis's extent +// out of a domain_traits-typed domain rather than taking them as separate arguments. +template +KMM_HOST_DEVICE make_strides_t make_strides_for_domain( + const DomainT& domain, + StrideT alignment, + IndexSequence +) { + Shape extents = { + checked_cast(domain_traits::extent(domain, Is))... + }; + return make_strides_from_shape(extents, alignment); +} + +} // namespace detail + +/// \addtogroup layout +/// @{ + +template +struct LayoutPolicy { + template + using mapping_type = make_strides_t< // + order, + nonvoid_t>, + domain_rank>; + + template + KMM_HOST_DEVICE mapping_type apply(const DomainT& domain) const noexcept { + using stride_type = nonvoid_t>; + + return detail::make_strides_for_domain( + domain, + checked_cast(alignment), + make_index_sequence>() + ); + } +}; + +/// A stride policy laying out a domain in column-major order (the first axis varies fastest +/// in memory). +/// +/// The leading axis is padded up to a multiple of `alignment` elements. +template +struct ColMajorPadded: LayoutPolicy {}; + +/// A stride policy laying out a domain in row-major order (the last axis varies fastest in +/// memory). +/// +/// The trailing axis padded up to a multiple of `alignment` elements. +template +struct RowMajorPadded: LayoutPolicy {}; + +/// Column-major stride policy (the first axis varies fastest in memory). +struct ColMajor: LayoutPolicy {}; + +/// Row-major stride policy (the last axis varies fastest in memory). +struct RowMajor: LayoutPolicy {}; + +/// A stride policy laying out a domain in either column-major order or row-major order, depending +/// on a provided runtime value. +/// +/// The contiguous axis is padded up to a multiple of `alignment` elements. +template +struct StridedPadded { + MemoryOrder order; + + KMM_HOST_DEVICE + StridedPadded(MemoryOrder order = MemoryOrder::RowMajor) : order(order) {} + + template + using mapping_type = StridesN< // + domain_index_type, + domain_rank>; + + template + KMM_HOST_DEVICE mapping_type apply(const DomainT& domain) const noexcept { + if (order == MemoryOrder::RowMajor) { + return RowMajorPadded {}.apply(domain); + } else { + return ColMajorPadded {}.apply(domain); + } + } +}; + +/// Policy that is either column-major or row-major major, depending on a runtime value. +struct Strided: StridedPadded<1> { + KMM_HOST_DEVICE + Strided(MemoryOrder order = MemoryOrder::RowMajor) : StridedPadded<1>(order) {} +}; + +/// Combines a domain (Shape or Bounds) with a stride mapping to describe how n-dimensional +/// indices translate to linear storage offsets. `PolicyT` is a stride policy (e.g. `RowMajor`, +/// `ColMajor`) from which the actual mapping is derived via `detail::policy_traits`; a mapping +/// (e.g. `Strides<...>`) can also be passed directly as `PolicyT`, since it acts as its own +/// (domain-independent) policy. Defaults to `RowMajor`. +template +class Layout { + public: + using self_type = Layout; + using domain_type = DomainT; + using policy_type = PolicyT; + + using domain_traits = detail::domain_traits; + static constexpr size_t rank = domain_traits::rank; + using index_type = typename domain_traits::index_type; + using ndindex_type = Point; + using shape_type = Shape; + using range_type = Range; + using bounds_type = Bounds; + + using mapping_traits = detail::mapping_traits; + using mapping_type = typename mapping_traits::mapping_type; + static_assert(mapping_traits::rank == rank, "domain and mapping must have the same rank"); + using stride_type = typename mapping_traits::stride_type; + using ndstrides_type = Vec; + + using zero_origin_type = Layout; + using move_origin_type = Layout; + + template + using drop_axis_type = Layout< // + typename domain_traits::template drop_axis_type, + typename mapping_traits::template drop_axis_type>; + + template + using insert_axis_type = Layout< // + typename domain_traits::template insert_axis_type, + typename mapping_traits::template insert_axis_type>; + + template + using permute_axes_type = Layout< // + typename domain_traits::template permute_axes_type, + typename mapping_traits::template permute_axes_type>; + + private: + template + struct permute_axes_seq_type_helper; + + template + struct permute_axes_seq_type_helper> { + using type = permute_axes_type; + }; + + public: + using reverse_axes_type = + typename permute_axes_seq_type_helper>::type; + + template + using swap_axes_type = + typename permute_axes_seq_type_helper>::type; + + using transpose_type = swap_axes_type<0, 1>; + + template + using move_axis_to_position_type = typename permute_axes_seq_type_helper< + move_axis_to_position_index_sequence>::type; + + template + using move_axis_to_front_type = move_axis_to_position_type; + + template + using move_axis_to_back_type = move_axis_to_position_type; + + template + using slice_axis_type = typename detail::slice_axis_impl::type; + + template + using slice_type = typename detail::slice_multi_impl::type; + + /// Constructs an empty (default-initialized domain and mapping) layout. + KMM_HOST_DEVICE + constexpr Layout() = default; + + /// Constructs a layout from a domain, a mapping, and an explicit base offset. + KMM_HOST_DEVICE + constexpr Layout(domain_type domain, mapping_type mapping, ptrdiff_t base_offset) : + m_domain(domain), + m_mapping(mapping), + m_base_offset(base_offset) {} + + /// Constructs a layout from a domain and a mapping, deriving the base offset from the origin. + KMM_HOST_DEVICE + constexpr Layout(domain_type domain, mapping_type mapping) : + m_domain(domain), + m_mapping(mapping) { + m_base_offset = -offset_span().start; + } + + /// Constructs a layout from a domain and a policy. The template argument is needed since + /// otherwise this constructor would be ambiguous with `Layout(domain_type, mapping_type)`. + template + KMM_HOST_DEVICE constexpr Layout(domain_type domain, P policy = {}) : + m_domain(domain), + m_mapping(detail::policy_traits::apply(policy, domain)) { + m_base_offset = -offset_span().start; + } + + /// Converting constructor from a layout over a compatible domain/mapping type. + template + KMM_HOST_DEVICE constexpr Layout(const Layout& that) : + m_domain(that.domain()), + m_mapping(that.mapping()), + m_base_offset(that.base_offset()) {} + + /// Returns the domain (shape or bounds) of this layout. + KMM_HOST_DEVICE + constexpr const domain_type& domain() const noexcept { + return m_domain; + } + + /// Returns the stride mapping of this layout. + KMM_HOST_DEVICE + constexpr const mapping_type& mapping() const noexcept { + return m_mapping; + } + + /// Returns the valid index range along the given axis, or a single-element range if out of rank. + KMM_HOST_DEVICE + range_type bounds(size_t axis) const noexcept { + return is_less(axis, rank) ? domain_traits::bounds(m_domain, axis) : range_type::one(); + } + + /// Returns the first valid index along the given axis. + KMM_HOST_DEVICE + index_type origin(size_t axis) const noexcept { + return begin(axis); + } + + /// Returns the first valid index along each axis. + KMM_HOST_DEVICE + ndindex_type origin() const noexcept { + return begin(); + } + + /// Returns the number of valid indices along the given axis, or 1 if out of rank. + KMM_HOST_DEVICE + index_type extent(size_t axis) const noexcept { + return is_less(axis, rank) ? domain_traits::extent(m_domain, axis) + : static_cast(1); + } + + /// Returns the start of the valid range along the given axis. + KMM_HOST_DEVICE + index_type begin(size_t axis) const noexcept { + return bounds(axis).start; + } + + /// Returns the end (exclusive) of the valid range along the given axis. + KMM_HOST_DEVICE + index_type end(size_t axis) const noexcept { + return bounds(axis).stop; + } + + /// Returns the start of the valid range along each axis. + KMM_HOST_DEVICE + ndindex_type begin() const noexcept { + ndindex_type result; + + for (size_t axis = 0; is_less(axis, rank); axis++) { + result[axis] = this->begin(axis); + } + + return result; + } + + /// Returns the end (exclusive) of the valid range along each axis. + KMM_HOST_DEVICE + ndindex_type end() const noexcept { + ndindex_type result; + + for (size_t axis = 0; is_less(axis, rank); axis++) { + result[axis] = this->end(axis); + } + + return result; + } + + /// Returns the extent (size) along each axis. + KMM_HOST_DEVICE + shape_type shape() const noexcept { + shape_type result; + + for (size_t axis = 0; is_less(axis, rank); axis++) { + result[axis] = this->extent(axis); + } + + return result; + } + + /// Returns the bounds (begin/end per axis) covered by this layout. + KMM_HOST_DEVICE + bounds_type bounds() const noexcept { + // from_bounds wants a Point, and Point's constructor from a Vec is explicit (no + // implicit conversion), so ndindex_type's Vec values need converting explicitly here. + return bounds_type::from_bounds( + Point(this->begin()), + Point(this->end()) + ); + } + + /// Returns the total number of elements covered by this layout. + KMM_HOST_DEVICE + index_type size() const noexcept { + return bounds().volume(); + } + + /// Returns whether this layout covers zero elements. + KMM_HOST_DEVICE + bool is_empty() const noexcept { + return bounds().is_empty(); + } + + /// Returns whether the given index falls within this layout's bounds. + KMM_HOST_DEVICE + bool contains(const ndindex_type& p) const noexcept { + return bounds().contains(Point(p)); + } + + /// Returns the stride along the given axis. + KMM_HOST_DEVICE + stride_type stride(size_t axis) const noexcept { + return mapping_traits::stride(m_mapping, axis); + } + + /// Returns the stride along each axis. + KMM_HOST_DEVICE + ndstrides_type strides() const noexcept { + ndstrides_type result; + + for (size_t axis = 0; is_less(axis, rank); axis++) { + result[axis] = this->stride(axis); + } + + return result; + } + + /// Returns the constant offset independent of any index, e.g. from a non-zero domain origin. + KMM_HOST_DEVICE + constexpr ptrdiff_t base_offset() const noexcept { + return m_base_offset; + } + + /// Computes the offset corresponding to an n-dimensional index, relative to `base_offset()`. + KMM_HOST_DEVICE + ptrdiff_t local_offset(const ndindex_type& index) const noexcept { + return mapping_traits::linearize_offset(m_mapping, index); + } + + /// Returns the total offset to add to the raw data pointer to reach `index`: + /// `base_offset() + local_offset(index)`. + KMM_HOST_DEVICE + ptrdiff_t offset(const ndindex_type& index) const noexcept { + return base_offset() + local_offset(index); + } + + /// Returns the range of numbers that can be returned from `offset(ndindex)`. In other words, + /// this gives the range of the lowest offset to the largest offset that will be accessed on + /// the storage that this layout applies to. + KMM_HOST_DEVICE + Range offset_span() const noexcept { + ptrdiff_t lo = m_base_offset; + ptrdiff_t hi = m_base_offset; + + if (!is_empty()) { + for (size_t i = 0; is_less(i, rank); i++) { + auto s = static_cast(stride(i)); + auto [a, b] = bounds(i); + lo += static_cast(s >= 0 ? a : b - 1) * s; + hi += static_cast(s >= 0 ? b : a - 1) * s; + } + } + + return {lo, hi}; + } + + /// Returns this layout with `delta` added to its base offset (domain and mapping unchanged). + KMM_HOST_DEVICE + constexpr self_type shift_offset(ptrdiff_t delta) const noexcept { + return self_type {m_domain, m_mapping, m_base_offset + delta}; + } + + /// Returns this layout shifted so `offset_range()` starts at zero (i.e., becomes `(0, N)`). + KMM_HOST_DEVICE + self_type normalize_offset() const noexcept { + return shift_offset(-offset_span().start); + } + + /// Returns whether the strides of this layout where produced by the given policy. + template + KMM_HOST_DEVICE bool is_mapping_from_policy(const OtherPolicyT& policy = {}) const noexcept { + return m_mapping + == detail::policy_traits::apply(policy, m_domain); + } + + /// Returns whether the layout's strides are contiguous in the given memory order. + KMM_HOST_DEVICE + bool is_contiguous(MemoryOrder order = MemoryOrder::RowMajor) const noexcept { + if (order == MemoryOrder::RowMajor) { + return m_mapping == make_strides_from_shape(shape()); + } else { + return m_mapping == make_strides_from_shape(shape()); + } + } + + /// Returns a copy of this layout with the domain replaced (mapping and base offset kept). + template + KMM_HOST_DEVICE Layout with_domain( + const NewDomainT& new_domain + ) const noexcept { + return Layout {new_domain, m_mapping, m_base_offset}; + } + + /// Returns a copy of this layout with the mapping replaced (domain and base offset kept). + template + KMM_HOST_DEVICE Layout with_mapping( + const NewMappingT& new_mapping + ) const noexcept { + return Layout {m_domain, new_mapping, m_base_offset}; + } + + /// Returns this layout rebased so its domain starts at the zero index. + KMM_HOST_DEVICE + zero_origin_type zero_origin() const noexcept { + return zero_origin_type {shape(), m_mapping, m_base_offset + local_offset(origin())}; + } + + /// Returns this layout shifted so it originates at the given index, keeping the same shape. + KMM_HOST_DEVICE + move_origin_type move_origin(ndindex_type new_origin) const noexcept { + auto new_domain = bounds_type::from_offset_size(Point::zero(), shape()); + auto offset_diff = local_offset(new_origin) - local_offset(origin()); + + return move_origin_type {new_domain, m_mapping, m_base_offset + offset_diff}; + } + + /// Returns this layout restricted to the intersection of its bounds and the given bounds. + KMM_HOST_DEVICE + move_origin_type restrict_bounds(bounds_type new_bounds) const noexcept { + return with_domain(new_bounds.intersection(bounds())); + } + + /// Returns this layout restricted along one axis to the intersection with [start, stop). + template + KMM_HOST_DEVICE move_origin_type + restrict_axis(index_type start, index_type stop) const noexcept { + static_assert(Axis < rank, "axis out of bounds"); + auto new_bounds = bounds(); + new_bounds[Axis] = range_type {start, stop}.intersection(new_bounds[Axis]); + return with_domain(new_bounds); + } + + /// Returns this layout sliced to the given bounds and rebased to a zero origin. + KMM_HOST_DEVICE + zero_origin_type slice_bounds(bounds_type new_bounds) const noexcept { + return with_domain(new_bounds).zero_origin(); + } + + /// Returns this layout with the given axis dropped, fixed at the given index. + template + KMM_HOST_DEVICE drop_axis_type drop_axis(index_type index) const noexcept { + static_assert(Axis < rank, "axis out of bounds"); + KMM_ASSERT(is_less(index, extent(Axis))); + + return drop_axis_type { + domain_traits::template drop_axis(domain()), + mapping_traits::template drop_axis(mapping()), + m_base_offset + static_cast(stride(Axis)) * static_cast(index) + }; + } + + /// Returns this layout with a new broadcast axis of the given extent inserted at the given position. + template + KMM_HOST_DEVICE insert_axis_type insert_axis( + index_type extent = static_cast(1) + ) const noexcept { + static_assert(Axis <= rank, "axis out of bounds"); + + return insert_axis_type { + domain_traits::template insert_axis(domain(), extent), + mapping_traits::template insert_axis(mapping()), + m_base_offset + }; + } + + /// Returns this layout with its axes reordered according to the given permutation, e.g. + /// `permute_axes<2, 0, 1>()` moves the current axis 2 to position 0, axis 0 to position 1, + /// and axis 1 to position 2. + template + KMM_HOST_DEVICE permute_axes_type permute_axes( + IndexSequence = {} + ) const noexcept { + static_assert(sizeof...(Is) == rank, "permutation must contain exactly `rank` axes"); + static_assert(is_permutation>, "must be a valid permutation of axes"); + + return permute_axes_type { + domain_traits::template permute_axes(domain(), IndexSequence {}), + mapping_traits::template permute_axes(mapping(), IndexSequence {}), + m_base_offset + }; + } + + /// Returns this layout with the order of all axes reversed. + KMM_HOST_DEVICE reverse_axes_type reverse_axes() const noexcept { + return permute_axes(reverse_index_sequence()); + } + + /// Returns this layout with axes `I` and `J` swapped. + template + KMM_HOST_DEVICE swap_axes_type swap_axes() const noexcept { + static_assert(I < rank && J < rank, "axis out of bounds"); + return permute_axes(swap_index_sequence()); + } + + /// Returns this layout with axes 0 and 1 swapped. Only valid for a rank-2 layout; use + /// `swap_axes` or `permute_axes` for other ranks. + KMM_HOST_DEVICE transpose_type transpose() const noexcept { + static_assert(rank == 2, "transpose() requires a rank-2 layout"); + return swap_axes<0, 1>(); + } + + /// Returns this layout with the given axis moved to the given position, preserving the + /// relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_position_type move_axis_to_position() const noexcept { + static_assert(Axis < rank && Pos < rank, "axis out of bounds"); + return permute_axes(move_axis_to_position_index_sequence()); + } + + /// Returns this layout with the given axis moved to the front (position 0), preserving the + /// relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_front_type move_axis_to_front() const noexcept { + return move_axis_to_position(); + } + + /// Returns this layout with the given axis moved to the back (position `rank - 1`), + /// preserving the relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_back_type move_axis_to_back() const noexcept { + return move_axis_to_position(); + } + + /// Returns this layout with the given axis sliced according to the given slice token (e.g. `all`, a `Range`, `new_axis`). + template + slice_axis_type slice_axis(const SliceT& slice) const noexcept { + return detail::slice_axis_impl::apply(*this, slice); + } + + /// Returns this layout with the given axis narrowed to the range [start, end). + template + KMM_HOST_DEVICE self_type slice_axis(index_type start, index_type end) const noexcept { + static_assert(Axis < rank, "axis out of bounds"); + + return self_type { + domain_traits::template slice_axis(domain(), start, end), + m_mapping, + m_base_offset + static_cast(stride(Axis)) * static_cast(start) + }; + } + + /// Returns this layout sliced across all axes at once, one slice token per axis. + template + slice_type slice(const Slices&... slices) const noexcept { + return detail::slice_multi_impl::apply(*this, slices...); + } + + private: + KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS domain_type m_domain {}; + KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS mapping_type m_mapping {}; + ptrdiff_t m_base_offset = 0; +}; + +/// @} + +/// \addtogroup layout +/// @{ + +/// Constructs a Layout for the given domain using the given stride policy. +template +KMM_HOST_DEVICE Layout make_layout(DomainT domain, PolicyT policy = {}) { + return Layout {domain, policy}; +} + +/// @} + +} // namespace kmm + +// Layout has no base class to delegate to, so this prints/hashes it as the tuple of its +// three constituent parts, each of which is itself printable/hashable (domain: Shape/Bounds, +// mapping: Strides, base_offset: a plain integer). +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +namespace kmm { +template +std::ostream& operator<<(std::ostream& stream, const Layout& layout) { + return stream << "Layout(domain=" << layout.bounds() << ", strides=" << layout.strides() + << ", base_offset=" << layout.base_offset() << ")"; +} +} // namespace kmm + +template +struct fmt::formatter>: fmt::ostream_formatter {}; +#endif diff --git a/include/kmm/utils/macros.hpp b/include/kmm/core/macros.hpp similarity index 82% rename from include/kmm/utils/macros.hpp rename to include/kmm/core/macros.hpp index cc6c8ae3..ef2e9694 100644 --- a/include/kmm/utils/macros.hpp +++ b/include/kmm/core/macros.hpp @@ -3,7 +3,6 @@ #define KMM_INLINE __attribute__((always_inline)) inline #define KMM_NOINLINE __attribute__((noinline)) -#define KMM_ASSUME(expr) (__builtin_assume((expr))) #define KMM_LIKELY(expr) (__builtin_expect(!!(expr), true)) #define KMM_UNLIKELY(expr) (__builtin_expect(!!(expr), false)) @@ -16,6 +15,7 @@ #define KMM_DEVICE __device__ __forceinline__ #define KMM_HOST_DEVICE_NOINLINE __host__ __device__ #define KMM_DEVICE_NOINLINE __device__ + #define KMM_LAMBDA [=] __device__ #ifdef __CUDA_ARCH__ #define KMM_IS_DEVICE (1) @@ -25,11 +25,12 @@ #endif #elif defined(__HIPCC__) // HIP - #include + // TODO: using __forceinline__ breaks compilation, it needs to be investigated #define KMM_HOST_DEVICE __host__ __device__ inline __attribute__((always_inline)) #define KMM_DEVICE __device__ inline __attribute__((always_inline)) #define KMM_HOST_DEVICE_NOINLINE __host__ __device__ #define KMM_DEVICE_NOINLINE __device__ + #define KMM_LAMBDA [=] __device__ #ifdef __HIP_DEVICE_COMPILE__ #define KMM_IS_DEVICE (1) @@ -43,10 +44,21 @@ #define KMM_DEVICE KMM_INLINE #define KMM_HOST_DEVICE_NOINLINE #define KMM_DEVICE_NOINLINE + #define KMM_LAMBDA [=] #define KMM_IS_RTC (0) #define KMM_IS_DEVICE (0) #endif +#ifdef __has_cpp_attribute + #if __has_cpp_attribute(no_unique_address) && !(defined(__GNUC__) && (__cplusplus < 201100)) + #define KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS [[no_unique_address]] + #endif +#endif + +#ifndef KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS + #define KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS +#endif + #define KMM_NOT_COPYABLE(TYPE) \ public: \ TYPE(const TYPE&) = delete; \ diff --git a/include/kmm/utils/panic.hpp b/include/kmm/core/panic.hpp similarity index 54% rename from include/kmm/utils/panic.hpp rename to include/kmm/core/panic.hpp index b2a9cd6d..55d4dcdc 100644 --- a/include/kmm/utils/panic.hpp +++ b/include/kmm/core/panic.hpp @@ -1,13 +1,11 @@ #pragma once -#include "macros.hpp" +#include "kmm/core/macros.hpp" #define KMM_PANIC(...) \ do { \ ::kmm::panic(__FILE__, __LINE__, __VA_ARGS__); \ - while (1) \ - ; \ - } while (0) + } while (1) #define KMM_ASSERT(...) \ do { \ @@ -19,7 +17,17 @@ #define KMM_DEBUG_ASSERT(...) KMM_ASSERT(__VA_ARGS__) #define KMM_TODO() KMM_PANIC("not implemented") -#if !KMM_IS_RTC +#ifndef NDEBUG + #define KMM_UNSAFE_ASSUME(expr) KMM_ASSERT(expr) +#else + #define KMM_UNSAFE_ASSUME(expr) \ + do { \ + if (!(expr)) \ + __builtin_unreachable(); \ + } while (0) +#endif + +#if !KMM_IS_DEVICE #include "fmt/format.h" #define KMM_PANIC_FMT(...) \ @@ -28,22 +36,33 @@ while (1) \ ; \ } while (0) +#else + // blockIdx/threadIdx are only declared once the platform's runtime header is included; unlike + // CUDA, HIP does not make them available as bare compiler builtins. This include must stay + // outside `namespace kmm` -- putting it inside nests everything the runtime header transitively + // pulls in (e.g. libstdc++'s -> ) under kmm::, which breaks std:: + // lookups throughout the translation unit. + #if defined(__CUDACC__) + #include + #elif defined(__HIPCC__) + #include + #endif #endif namespace kmm { +/// \addtogroup utility +/// @{ + #if !KMM_IS_DEVICE -/** - * Logs a fatal error, prints relevant debugging info, and aborts the program. - * - * @param file Source filename where the panic occurred. - * @param line Line number where the panic occurred. - * @param function Function name where the panic occurred. - * @param message Reason for the panic. - */ +/// Logs a fatal error, prints relevant debugging info, and aborts the program. +/// +/// `file` and `line` are the source location where the panic occurred, and `message` is the +/// reason for the panic. [[noreturn]] void panic(const char* filename, int lineno, const char* message); #else -KMM_DEVICE void panic(const char* filename, int lineno, const char* message) { + +[[noreturn]] KMM_DEVICE void panic(const char* filename, int lineno, const char* message) { printf( "[block=(%u,%u,%u) thread=(%u,%u,%u)] PANIC at %s:%d: %s\n", blockIdx.x, @@ -58,9 +77,15 @@ KMM_DEVICE void panic(const char* filename, int lineno, const char* message) { ); while (true) { + #if defined(__HIPCC__) + __builtin_trap(); + #else asm volatile("trap;"); + #endif } } #endif -} // namespace kmm \ No newline at end of file +/// @} + +} // namespace kmm diff --git a/include/kmm/core/point.hpp b/include/kmm/core/point.hpp new file mode 100644 index 00000000..28672f3f --- /dev/null +++ b/include/kmm/core/point.hpp @@ -0,0 +1,154 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/vec.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// An N-dimensional coordinate, backed by a `Vec`. +/// +/// A `Point` identifies a location within an N-dimensional index space (e.g. the position +/// of an element in an N-dimensional array). +template +class Point: public Vec { + public: + using storage_type = Vec; + + /// Construct point from vector. + KMM_HOST_DEVICE + constexpr Point(const storage_type& storage) : storage_type(storage) {} + + /// Construct point (0, 0, 0, ...). + KMM_HOST_DEVICE + constexpr Point() : storage_type(fill(static_cast(0))) {} + + /// Constructs a point from N values. + template> + KMM_HOST_DEVICE Point(T first, Ts&&... args) : storage_type {first, args...} {} + + /// Converts from a point of a different dimensionality/type, throwing on overflow. + template + KMM_HOST_DEVICE constexpr Point(const Point& that) : storage_type(Point::from(that)) { + if (!that.template is_convertible_to()) { + throw_overflow_exception(); + } + } + + /// Builds a point from a vector, padding any missing dimensions with zero. + template + KMM_HOST_DEVICE static constexpr Point from(const Vec& that) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = is_less(i, M) ? static_cast(that[i]) : static_cast(0); + } + + return Point(result); + } + + /// Creates a point with every coordinate set to one. + KMM_HOST_DEVICE + static constexpr Point one() { + return make_index_sequence::template fill(static_cast(1)); + } + + /// Creates a point with every coordinate set to zero. + KMM_HOST_DEVICE + static constexpr Point zero() { + return make_index_sequence::template fill(static_cast(0)); + } + + /// Checks whether this point fits losslessly into a Point. + template + KMM_HOST_DEVICE bool is_convertible_to() const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + if (is_less(i, M)) { + result &= is_convertible((*this)[i]); + } else { + result &= is_equal((*this)[i], static_cast(0)); + } + } + + return result; + } + + /// Returns coordinate i, or default_value if the axis is out of range. + KMM_HOST_DEVICE + T get_or_default(size_t i, T default_value = T {}) const { + if constexpr (N > 0) { + if (KMM_LIKELY(is_less(i, N))) { + return (*this)[i]; + } + } + + return default_value; + } +}; + +template +Point(Ts&&...) -> Point; + +/// Constructs a Point from the given coordinate values. +template +KMM_HOST_DEVICE Point point(const Ts&... values) { + return Point {Vec {values...}}; +} + +template +KMM_HOST_DEVICE Point concat(const Point& lhs, const Point& rhs) { + return Point { + concat(static_cast&>(lhs), static_cast&>(rhs)) + }; +} + +template +KMM_HOST_DEVICE bool operator==(const Point& lhs, const Point& rhs) { + bool result = true; + + for (size_t i = 0; is_less(i, N) || is_less(i, M); i++) { + result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); + } + + return result; +} + +template +KMM_HOST_DEVICE bool operator!=(const Point& lhs, const Point& rhs) { + return !(lhs == rhs); +} + +/// Adds two points coordinate-wise. +template +KMM_HOST_DEVICE constexpr Point operator+(const Point& lhs, const Point& rhs) { + Point result = lhs; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] += rhs[i]; + } + + return result; +} + +/// @} + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash>: std::hash> {}; +#endif \ No newline at end of file diff --git a/include/kmm/core/range.hpp b/include/kmm/core/range.hpp new file mode 100644 index 00000000..247444cd --- /dev/null +++ b/include/kmm/core/range.hpp @@ -0,0 +1,299 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// Represents a half-open range of numbers beginning at `start` and ending at `stop`, excluding +/// the `stop` value itself. For example, `Range(5, 10)` represents the numbers +/// `5, 6, 7, 8, 9`, but does not include `10`. +/// +/// A range with `start < stop` is considered valid and non-empty (size is `stop-start`). +/// A range with `start == stop` is considered valid but empty (size is 0). +/// A range with `start > stop` is considered invalid (most operations treat this as empty). +/// +/// Ranges can be iterated over directly: +/// +/// ``` +/// for (auto i : Range(5, 10)) { +/// printf("%d\n", i); // prints 5, 6, 7, 8, 9 +/// } +/// ``` +template +class Range { + public: + struct iterator; + using value_type = T; + using const_iterator = iterator; + + constexpr Range(const Range&) = default; + constexpr Range(Range&&) noexcept = default; + + constexpr Range& operator=(const Range&) = default; + constexpr Range& operator=(Range&&) noexcept = default; + + /// Constructs an empty range `0...0` + KMM_HOST_DEVICE + constexpr Range() = default; + + /// Constructs the range `0...stop` + KMM_HOST_DEVICE + explicit constexpr Range(T stop) : stop(stop) {} + + /// Constructs the range `start...stop` + KMM_HOST_DEVICE + constexpr Range(T start, T stop) : start(start), stop(stop) {} + + /// Converts a range from another range. Throws exception on overflow. + template + KMM_HOST_DEVICE constexpr Range(const Range& that) { + if (!that.template is_convertible_to()) { + throw_overflow_exception(); + } + + *this = Range::from(that); + } + + /// Converts a range from another range. Does not check for overflow. + template + KMM_HOST_DEVICE static constexpr Range from(const Range& range) { + return {static_cast(range.start), static_cast(range.stop)}; + } + + /// Returns the range 0...1 + KMM_HOST_DEVICE static constexpr Range one() { + return {static_cast(0), static_cast(1)}; + } + + /// Returns whether both bounds can be converted to `U` without overflow. + template + KMM_HOST_DEVICE constexpr bool is_convertible_to() const { + return is_convertible(start) && is_convertible(stop); + } + + /// Checks if the range is empty (i.e., `begin == end`) or invalid (i.e., `begin > end`). + KMM_HOST_DEVICE + constexpr bool is_empty() const noexcept { + return !(this->start < this->stop); + } + + /// Returns an iterator to the first element. + KMM_HOST_DEVICE + constexpr iterator begin() const noexcept { + return this->start; + } + + /// Returns an iterator to one passed the last element. + KMM_HOST_DEVICE + constexpr iterator end() const noexcept { + return is_empty() ? this->start : this->stop; + } + + // Checks if the given index `index` is within this range. + template + KMM_HOST_DEVICE constexpr bool contains(const U& index) const noexcept { + return is_less_equal(this->start, index) & is_less(index, this->stop); + } + + /// Returns whether `that` is fully contained within this range. + /// Empty ranges are considered contained. Invalid ranges are never contained. + template + KMM_HOST_DEVICE constexpr bool contains(const Range& that) const noexcept { + return is_less_equal(this->start, that.start) & // + is_less_equal(that.stop, this->stop) & // + is_less_equal(this->start, this->stop) & // + is_less_equal(that.start, that.stop); + } + + /// Returns whether this range and `that` overlap with non-empty intersection. + /// Empty or invalid ranges never overlap. + template + KMM_HOST_DEVICE constexpr bool overlaps(const Range& that) const noexcept { + return is_less(this->start, this->stop) & // + is_less(that.start, that.stop) & // + is_less(this->start, that.stop) & // + is_less(that.start, this->stop); + } + + /// Returns the range that lies in the intersection of `this` and `that`. + KMM_HOST_DEVICE + constexpr Range intersection(const Range& that) const noexcept { + return { + this->start > that.start ? this->start : that.start, + this->stop < that.stop ? this->stop : that.stop, + }; + } + + /// Computes the size (or length) of the range. This always a non-negative number as + /// empty or invalid ranges will have a length of zero. + KMM_HOST_DEVICE + constexpr T size() const noexcept { + return this->start >= this->stop ? static_cast(0) : this->stop - this->start; + } + + struct Pair { + Range first; + Range second; + }; + + /// Returns the ranges `start..mid` and `mid..stop`. Example + /// + /// ``` + /// auto [left, right] = some_range.split(10); + /// ``` + KMM_HOST_DEVICE + constexpr Pair split(T mid) const { + if (mid < this->start) { + mid = this->start; + } + + if (mid > this->stop) { + mid = this->stop; + } + + return {{start, mid}, {mid, stop}}; + } + + /// The start point of the range. + T start = static_cast(0); + + /// The end point of the range (not inclusive). + T stop = static_cast(0); +}; + +template +Range(const T&) -> Range; + +template +Range(const T&, const T&) -> Range; + +/// Constructs the range `0...stop`. +template +KMM_HOST_DEVICE constexpr Range range(const T& stop) { + return Range(static_cast(0), stop); +} + +/// Constructs the range `start...stop`. +template +KMM_HOST_DEVICE constexpr Range range(const T& start, const T& stop) { + return Range(start, stop); +} + +template +KMM_HOST_DEVICE constexpr bool operator==(const Range& lhs, const Range& rhs) { + return is_equal(lhs.start, rhs.start) && is_equal(lhs.stop, rhs.stop); +} + +template +KMM_HOST_DEVICE constexpr bool operator!=(const Range& lhs, const Range& rhs) { + return !(lhs == rhs); +} + +/// Returns `Range(range.start + offset, range.stop + offset)` +template +KMM_HOST_DEVICE constexpr Range operator+(const Range& range, const T& offset) { + return {range.start + offset, range.stop + offset}; +} + +/// Returns `Range(range.start + offset, range.stop + offset)` +template +KMM_HOST_DEVICE constexpr Range operator+(const T& offset, const Range& range) { + return {offset + range.start, offset + range.stop}; +} + +/// Returns `Range(range.start - offset, range.stop - offset)` +template +KMM_HOST_DEVICE constexpr Range operator-(const Range& range, const T& offset) { + return {range.start - offset, range.stop - offset}; +} + +template +struct Range::iterator { + public: + using value_type = T; + using difference_type = + decltype(static_cast(nullptr) - static_cast(nullptr)); //std::ptrdiff_t; + using pointer = const T*; + using reference = const T&; + + KMM_HOST_DEVICE + constexpr iterator(T value) : current(value) {} + + KMM_HOST_DEVICE + constexpr const T& get() const { + return this->current; + } + + KMM_HOST_DEVICE + constexpr T& get() { + return this->current; + } + + KMM_HOST_DEVICE + constexpr operator T() const { + return this->current; + } + + KMM_HOST_DEVICE + constexpr T operator*() const { + return this->current; + } + + KMM_HOST_DEVICE + constexpr iterator& operator++() { + ++this->current; + return *this; + } + + KMM_HOST_DEVICE + constexpr iterator operator++(int) { + return iterator(this->current++); + } + + KMM_HOST_DEVICE + friend constexpr bool operator==(const iterator& lhs, const iterator& rhs) { + return lhs.current == rhs.current; + } + + KMM_HOST_DEVICE + friend constexpr bool operator!=(const iterator& lhs, const iterator& rhs) { + return lhs.current != rhs.current; + } + + T current; +}; + +/// @} + +} // namespace kmm + +#if !KMM_IS_RTC + #include + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +namespace kmm { + +template +std::ostream& operator<<(std::ostream& stream, const Range& p) { + return stream << p.start << "..." << p.stop; +} + +} // namespace kmm + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash> { + size_t operator()(const kmm::Range& p) const { + return ::kmm::hash_fields(p.start, p.stop); + } +}; +#endif \ No newline at end of file diff --git a/include/kmm/core/reduction.hpp b/include/kmm/core/reduction.hpp deleted file mode 100644 index 75e54963..00000000 --- a/include/kmm/core/reduction.hpp +++ /dev/null @@ -1,38 +0,0 @@ -#pragma once - -#include "kmm/core/data_type.hpp" -#include "kmm/core/identifiers.hpp" - -namespace kmm { - -enum struct Reduction : uint8_t { Invalid = 0, Sum, Product, Min, Max, BitAnd, BitOr }; - -std::vector reduction_identity_value(DataType dtype, Reduction op); - -struct ReductionInput { - BufferId buffer_id; - MemoryId memory_id; - EventList dependencies; - size_t num_inputs_per_output = 1; -}; - -struct ReductionOutput { - Reduction operation; - DataType data_type; - size_t num_outputs; -}; - -std::ostream& operator<<(std::ostream& f, Reduction p); -std::ostream& operator<<(std::ostream& f, ReductionInput p); -std::ostream& operator<<(std::ostream& f, ReductionOutput p); - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/core/resource.hpp b/include/kmm/core/resource.hpp deleted file mode 100644 index 47f3c73a..00000000 --- a/include/kmm/core/resource.hpp +++ /dev/null @@ -1,254 +0,0 @@ -#pragma once - -#include "kmm/core/buffer.hpp" -#include "kmm/core/system_info.hpp" -#include "kmm/core/view.hpp" -#include "kmm/utils/checked_math.hpp" -#include "kmm/utils/gpu_utils.hpp" - -namespace kmm { - -class Resource; -class InvalidResourceException; -class ComputeTask; -struct TaskContext; - -enum struct ExecutionSpace { Host, Device }; - -struct TaskContext { - std::vector accessors; -}; - -class ComputeTask { - public: - virtual ~ComputeTask() = default; - virtual void execute(Resource& resource, TaskContext context) = 0; -}; - -/** - * Exception throw if invalid resource is provided to task. - */ -class InvalidResourceException: public std::exception { - public: - InvalidResourceException(const std::type_info& expected, const std::type_info& gotten); - const char* what() const noexcept override; - - private: - std::string m_message; -}; - -class Resource { - public: - virtual ~Resource() = default; - - template - T* cast_if() noexcept { - return dynamic_cast(this); - } - - template - const T* cast_if() const noexcept { - return dynamic_cast(this); - } - - template - T& cast() { - if (auto* ptr = this->template cast_if()) { - return *ptr; - } - - throw InvalidResourceException(typeid(T), typeid(*this)); - } - - template - const T& cast() const { - if (auto* ptr = this->template cast_if()) { - return *ptr; - } - - throw InvalidResourceException(typeid(T), typeid(*this)); - } - - template - bool is() const noexcept { - return this->template cast_if() != nullptr; - } -}; - -class HostResource: public Resource {}; - -class DeviceResource: public DeviceInfo, public Resource { - KMM_NOT_COPYABLE_OR_MOVABLE(DeviceResource); - - public: - DeviceResource(DeviceInfo info, GPUContextHandle context, g_stream_t stream); - ~DeviceResource(); - - /** - * Returns a handle to the context associated with this device. - */ - GPUContextHandle context_handle() const { - return m_context; - } - - /** - * Returns a handle to the context associated with this device. - */ - g_context_t context() const { - return m_context; - } - - /** - * Returns a handle to the stream associated with this device. - */ - g_stream_t stream() const { - return m_stream; - } - - /** - * Shorthand for `stream()`. - */ - operator g_stream_t() const { - return m_stream; - } - - /** - * Returns a handle to the BLAS instance associated with this device. - */ - blas_handle_t blas() const { - return m_blas_handle; - } - - /** - * Block the current thread until all work submitted onto the stream of this device has - * finished. Note that the executor thread will also synchronize the stream automatically - * after each task, so calling thus function manually is not mandatory. - */ - void synchronize() const; - - /** - * Launch the given kernel onto the stream of this device. The `kernel_function` argument - * should be a pointer to a `__global__` function. - */ - template - void launch_impl( - dim3 grid_dim, - dim3 block_dim, - unsigned int shared_mem, - void (*const kernel_function)(Args...), - Args... args - ) const { - // Get void pointer to the arguments. - void* void_args[sizeof...(Args) + 1] = {static_cast(&args)..., nullptr}; - - // Launch the kernel! - // NOTE: This must be in the header file since `gpuLaunchKernel` seems to no find the - // kernel function if it is called from within a C++ file. - KMM_GPU_CHECK(gpu_launch_kernel( - reinterpret_cast(kernel_function), - grid_dim, - block_dim, - void_args, - shared_mem, - m_stream - )); - } - - template - void launch( - dim3 grid_dim, - dim3 block_dim, - unsigned int shared_mem, - void (*const kernel_function)(Param...), - Args... args - ) const { - launch_impl(grid_dim, block_dim, shared_mem, kernel_function, Param(args)...); - } - - /** - * Fill the provided buffer with the copies of the provided value. The fill is performed - * asynchronously on the stream of this device. - */ - template - void fill(T* dest, I num_elements, T value) const { - fill_bytes( - dest, - checked_mul(checked_cast(num_elements), sizeof(T)), - &value, - sizeof(T) - ); - } - - /** - * Fill the provided view with the copies of the provided value. The fill is performed - * asynchronously on the stream of this device. - */ - template - void fill(GPUViewMut dest, T value) const { - KMM_ASSERT(dest.is_contiguous()); - fill_bytes(dest.data(), dest.size_in_bytes(), &value, sizeof(T)); - } - - /** - * Copy data from the given source view to the given destination view. The copy is performed - * asynchronously on the stream of the current device. - */ - template - void copy(GPUView source, GPUViewMut dest) const { - KMM_ASSERT(source.sizes() == dest.sizes()); - KMM_ASSERT(source.is_contiguous() && dest.is_contiguous()); - copy_bytes(source.data(), dest.data(), source.size_in_bytes()); - } - - template - void copy(GPUView source, ViewMut dest) const { - KMM_ASSERT(source.sizes() == dest.sizes()); - KMM_ASSERT(source.is_contiguous() && dest.is_contiguous()); - copy_bytes(source.data(), dest.data(), source.size_in_bytes()); - } - - template - void copy(View source, GPUViewMut dest) const { - KMM_ASSERT(source.sizes() == dest.sizes()); - KMM_ASSERT(source.is_contiguous() && dest.is_contiguous()); - copy_bytes(source.data(), dest.data(), source.size_in_bytes()); - } - - /** - * Copy data from the given source view to the given destination view. The copy is performed - * asynchronously on the stream of the current device. - */ - template - void copy(const T* source_ptr, T* dest_ptr, I num_elements) const { - copy_bytes( - source_ptr, - dest_ptr, - checked_mul(checked_cast(num_elements), sizeof(T)) - ); - } - - /** - * Fill `nbytes` of the buffer starting at `dest_buffer` by repeating the given pattern. - * The argument `dest_buffer` must be allocated on the device while the `fill_pattern` must - * be on the host. - */ - void fill_bytes( - void* dest_buffer, - size_t nbytes, - const void* fill_pattern, - size_t fill_pattern_size - ) const; - - /** - * Copy `nbytes` bytes from the buffer starting at `source_buffer` to the buffer starting at - * `dest_buffer`. Both buffers must be allocated on the current device. - */ - void copy_bytes(const void* source_buffer, void* dest_buffer, size_t nbytes) const; - - private: - GPUContextHandle m_context; - g_stream_t m_stream; - blas_handle_t m_blas_handle; -}; - -} // namespace kmm diff --git a/include/kmm/core/shape.hpp b/include/kmm/core/shape.hpp new file mode 100644 index 00000000..fe7f22d2 --- /dev/null +++ b/include/kmm/core/shape.hpp @@ -0,0 +1,308 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/domain_traits.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/range.hpp" +#include "kmm/core/type_utils.hpp" +#include "kmm/core/vec.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// The extent (size) of an N-dimensional domain along each axis, backed by a `Vec`. +/// +/// A `Shape` describes how many elements exist along each axis of a domain (e.g. the +/// dimensions of an array). +template +class Shape: public Vec { + public: + using storage_type = Vec; + + KMM_HOST_DEVICE + explicit constexpr Shape(const storage_type& storage) : storage_type(storage) {} + + /// Create an empty shape `(0, 0, ...)` + KMM_HOST_DEVICE + constexpr Shape() : storage_type(::kmm::fill(static_cast(0))) {} + + constexpr Shape(const Shape&) = default; + constexpr Shape(Shape&&) noexcept = default; + Shape& operator=(const Shape&) = default; + Shape& operator=(Shape&&) noexcept = default; + + /// Create a shape `(first, args...)` + template> + KMM_HOST_DEVICE Shape(T first, Ts&&... args) : storage_type {first, args...} {} + + /// Create a shape from another shape. Throws on overflow. + template + KMM_HOST_DEVICE constexpr Shape(const Shape& that) : storage_type(Shape::from(that)) { + if (!that.template is_convertible_to()) { + throw_overflow_exception(); + } + } + + /// Create a shape from another shape. Does not throw on overflow. + template + KMM_HOST_DEVICE static constexpr Shape from(const Vec& that) { + storage_type result; + + for (size_t i = 0; is_less(i, N); i++) { + result[i] = is_less(i, M) ? static_cast(that[i]) : static_cast(1); + } + + return Shape(result); + } + + /// Create a shape `(value, value, value, ...)`. + KMM_HOST_DEVICE + static constexpr Shape fill(T value) { + return Shape(::kmm::fill(value)); + } + + /// Create a shape `(1, 1, 1, ...)`. + KMM_HOST_DEVICE + static constexpr Shape one() { + return fill(static_cast(1)); + } + + /// Create a shape `(0, 0, 0, ...)`. + KMM_HOST_DEVICE + static constexpr Shape zero() { + return fill(static_cast(0)); + } + + /// Returns coordinate i, or default_value if the axis is out of range. + KMM_HOST_DEVICE + T get_or_default(size_t i, T default_value = static_cast(1)) const { + if constexpr (N > 0) { + if (KMM_LIKELY(is_less(i, N))) { + return (*this)[i]; + } + } + + return default_value; + } + + /// Checks whether this shape can be converted to `Shape`. + template + KMM_HOST_DEVICE bool is_convertible_to() const { + bool result = true; + + for (size_t i = 0; is_less(i, N); i++) { + if (i < M) { + result &= is_convertible((*this)[i]); + } else { + result &= is_equal((*this)[i], static_cast(1)); + } + } + + return result; + } + + /// Check if this shape is empty. + /// + /// A shape is empty if each axis is less than or equal to zero. + KMM_HOST_DEVICE + bool is_empty() const { + bool result = false; + + for (size_t i = 0; is_less(i, N); i++) { + result |= !(static_cast(0) < (*this)[i]); + } + + return result; + } + + /// Returns the product of the extents of this shape. + /// + /// If one of the axis is negative, the returned value is zeor. + KMM_HOST_DEVICE + T volume() const { + T result = static_cast(1); + + if constexpr (N >= 1) { + result = (*this)[0]; + + for (size_t i = 1; is_less(i, N); i++) { + result *= (*this)[i]; + } + } + + return is_empty() ? static_cast(0) : result; + } + + /// Check if a point falls within this shape. + /// + /// This means that for each axis `p[i] >= 0` and `p[i] < shape[i]`. + template + KMM_HOST_DEVICE bool contains(const Point& p) const { + bool result = true; + + for (size_t i = 0; is_less(i, N) && is_less(i, M); i++) { + result &= !is_less(p[i], static_cast(0)) && is_less(p[i], (*this)[i]); + } + + if constexpr (N < M) { + for (size_t i = N; is_less(i, M); i++) { + result &= is_equal(p[i], static_cast(0)); + } + } + + if constexpr (N > M) { + for (size_t i = M; is_less(i, N); i++) { + result &= is_less(static_cast(0), (*this)[i]); + } + } + + return result; + } +}; + +template +Shape(Ts&&...) -> Shape; + +/// Constructs a Shape from the given per-axis extents. +template +KMM_HOST_DEVICE Shape shape(const Ts&... values) { + return Shape { + Vec {static_cast(values)...} + }; +} + +template +KMM_HOST_DEVICE Shape concat(const Shape& lhs, const Shape& rhs) { + return Shape { + concat(static_cast&>(lhs), static_cast&>(rhs)) + }; +} + +template +KMM_HOST_DEVICE bool operator==(const Shape& lhs, const Shape& rhs) { + bool result = true; + + for (size_t i = 0; is_less(i, N) || is_less(i, M); i++) { + result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); + } + + return result; +} + +template +KMM_HOST_DEVICE bool operator!=(const Shape& lhs, const Shape& rhs) { + return !(lhs == rhs); +} + +/// @} + +namespace detail { + +template +struct domain_traits> { + static constexpr size_t rank = N; + using index_type = IndexT; + using domain_type = Shape; + + KMM_HOST_DEVICE + static constexpr Range bounds(const domain_type& domain, size_t axis) { + return {static_cast(0), domain[axis]}; + } + + KMM_HOST_DEVICE + static constexpr index_type extent(const domain_type& domain, size_t axis) { + return domain[axis]; + } + + template + using slice_axis_type = Shape; + + template + KMM_HOST_DEVICE static constexpr slice_axis_type slice_axis( + domain_type domain, + index_type begin, + index_type end + ) { + domain[Axis] = end - begin; + return domain; + } + template + using drop_axis_type = Shape; + + template + KMM_HOST_DEVICE static constexpr drop_axis_type drop_axis(const domain_type& domain) { + return permute_axes(domain, drop_index_sequence()); + } + + /// Runtime-axis counterpart of `drop_axis`. The result type only depends on `N` (not on + /// which axis is dropped), so `axis` does not need to be known at compile time here. + KMM_HOST_DEVICE + static constexpr drop_axis_type<0> drop_axis(const domain_type& domain, size_t axis) { + KMM_ASSERT(axis < N); + drop_axis_type<0> result; + size_t j = 0; + + for (size_t i = 0; is_less(i, N); i++) { + if (i != axis) { + result[j] = domain[i]; + j++; + } + } + + return result; + } + + template + using permute_axes_type = Shape; + + template + KMM_HOST_DEVICE static constexpr permute_axes_type permute_axes( + const domain_type& domain, + IndexSequence + ) { + return {domain[Is]...}; + } + + template + using insert_axis_type = Shape; + + template + KMM_HOST_DEVICE static constexpr insert_axis_type insert_axis( + const domain_type& domain, + index_type extent + ) { + insert_axis_type result; + + for (size_t i = 0; is_less(i, Axis); i++) { + result[i] = domain[i]; + } + + result[Axis] = extent; + + for (size_t i = Axis; is_less(i, N); i++) { + result[i + 1] = domain[i]; + } + + return result; + } +}; + +} // namespace detail + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash>: std::hash> {}; +#endif \ No newline at end of file diff --git a/include/kmm/core/strides.hpp b/include/kmm/core/strides.hpp new file mode 100644 index 00000000..30798b4e --- /dev/null +++ b/include/kmm/core/strides.hpp @@ -0,0 +1,354 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/const_value.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/shape.hpp" +#include "kmm/core/type_utils.hpp" +#include "kmm/core/vec.hpp" + +namespace kmm { + +namespace detail { + +template +struct pack_element; + +template +struct pack_element: pack_element {}; + +template +struct pack_element<0, T, Rest...> { + using type = T; +}; + +template +struct strides_storage_type; + +template +struct strides_storage_type> { + static constexpr size_t rank = 0; + + template + using axis_stride_type = ConstValue; + + KMM_HOST_DEVICE + constexpr StrideT get_dynamic(size_t axis) const noexcept { + return 0; + } + + template + KMM_HOST_DEVICE constexpr axis_stride_type get_static(ConstIndex) const noexcept { + return {}; + } +}; + +template +struct strides_storage_type> { + using storage_type = StrideT; + static constexpr size_t rank = sizeof...(StridesT); + + template + struct convert_stride_impl { + KMM_HOST_DEVICE + static storage_type pack(const T& value) { + return static_cast(value); + } + + KMM_HOST_DEVICE + static T unpack(const storage_type& value) { + return static_cast(value); + } + }; + + template + struct convert_stride_impl> { + KMM_HOST_DEVICE + static storage_type pack(ConstValue) { + return static_cast(Value); + } + + KMM_HOST_DEVICE + static ConstValue unpack(storage_type) { + return {}; + } + }; + + template + using axis_stride_type = + typename pack_element>::type; + + KMM_HOST_DEVICE + constexpr strides_storage_type(StridesT... strides) noexcept : + m_values {convert_stride_impl::pack(strides)...} {} + + KMM_HOST_DEVICE + constexpr storage_type get_dynamic(size_t axis) const noexcept { + return axis < rank ? m_values[axis] : static_cast(0); + } + + template + KMM_HOST_DEVICE constexpr axis_stride_type get_static(ConstIndex) const noexcept { + return convert_stride_impl>::unpack(m_values[Axis]); + } + + storage_type m_values[rank]; +}; + +template +struct strides_storage_type...>> { + static constexpr size_t rank = sizeof...(Values); + + template + using axis_stride_type = + typename pack_element..., ConstValue>::type; + + KMM_HOST_DEVICE + constexpr strides_storage_type(ConstValue... strides) noexcept {} + + KMM_HOST_DEVICE + constexpr StrideT get_dynamic(size_t axis) const noexcept { + constexpr StrideT static_values[rank] = {static_cast(Values)...}; + return axis < rank ? static_values[axis] : static_cast(0); + } + + template + KMM_HOST_DEVICE constexpr axis_stride_type get_static(ConstIndex) const noexcept { + return {}; + } +}; + +} // namespace detail + +template +class Strides: + private detail::strides_storage_type> { + using base_type = detail::strides_storage_type>; + + public: + using stride_type = default_stride_type; + static constexpr size_t rank = sizeof...(StridesT); + + template + using axis_stride_type = typename base_type::template axis_stride_type; + + /// Default constructor. + template 0)>> + KMM_HOST_DEVICE constexpr Strides() : base_type(StridesT {}...) {} + + /// Construct from the given strides. + KMM_HOST_DEVICE + constexpr Strides(StridesT... strides) : base_type(strides...) {} + + /// Construct from another strides object. Throws exception on mismatched types. + template> + KMM_HOST_DEVICE constexpr Strides(const Strides& that) : + Strides(that, make_index_sequence()) {} + + /// Returns the `i`-th stride converted to `stride_type`. + KMM_HOST_DEVICE + constexpr stride_type operator[](size_t axis) const noexcept { + return base_type::get_dynamic(axis); + } + + /// Returns the `i`-th stride as + template + KMM_HOST_DEVICE constexpr axis_stride_type get( + ConstIndex index = {} + ) const noexcept { + return base_type::get_static(index); + } + + /// Returns the product `point[0] * strides[0] + point[1] * strides[1] + ...`. + template + KMM_HOST_DEVICE constexpr ptrdiff_t linearize_offset( + const Vec& point + ) const noexcept { + return linearize_offset_impl(point, make_index_sequence()); + } + + /// Construct from another strides object. Throws exception on mismatched types. + template + KMM_HOST_DEVICE bool is_equal(const Strides& that) const noexcept { + static constexpr size_t N = sizeof...(StridesT) > sizeof...(OtherStridesT) + ? sizeof...(StridesT) + : sizeof...(OtherStridesT); + return is_equal_impl(that, make_index_sequence()); + } + + KMM_HOST_DEVICE Vec to_vec() const noexcept { + return to_vec_impl(make_index_sequence()); + } + + private: + template + KMM_HOST_DEVICE constexpr Strides(const Strides& that, IndexSequence) : + base_type(checked_cast(that.get(ConstIndex {}))...) { + // Only reached via the public converting constructor, which is constrained to + // sizeof...(OtherStridesT) == rank, so `that` always has exactly `rank` axes here. + } + + template + KMM_HOST_DEVICE constexpr ptrdiff_t linearize_offset_impl( + const Vec& point, + IndexSequence + ) const noexcept { + return ( + static_cast(0) + ... + + (static_cast(point[Is]) + * static_cast(base_type::get_static(ConstIndex()))) + ); + } + + template + KMM_HOST_DEVICE bool is_equal_impl( + const Strides& that, + IndexSequence + ) const noexcept { + return (::kmm::is_equal(this->get(ConstIndex()), that.get(ConstIndex())) && ...); + } + + template + KMM_HOST_DEVICE Vec to_vec_impl(IndexSequence) const noexcept { + return {(*this)[Is]...}; + } +}; + +template +KMM_HOST_DEVICE bool operator==( + const Strides& lhs, + const Strides& rhs +) { + return lhs.is_equal(rhs); +} + +template +KMM_HOST_DEVICE bool operator!=( + const Strides& lhs, + const Strides& rhs +) { + return !(lhs == rhs); +} + +template +Strides make_strides(const StridesT&... strides) { + return {strides...}; +} + +namespace detail { +template> +struct strides_repeat_impl; + +template +struct strides_repeat_impl> { + // identity_t is needed just to expand the index sequence. + using type = Strides...>; +}; +} // namespace detail + +/// Alias for `Strides`, where the argument is repeated `N` times. +template +using StridesN = typename detail::strides_repeat_impl::type; + +/// Indicates whether the strides are ordered from from the trailing to the leading axis (RowMajor) +/// or from the leading to the trailing axis (ColMajor). +enum struct MemoryOrder { RowMajor, ColMajor }; + +/// @} + +namespace detail { + +template +struct calculate_product; + +template +struct calculate_product> { + using type = IndexT; + + template + KMM_HOST_DEVICE static type apply(const Shape& shape) { + return (static_cast(shape[Is]) * ... * static_cast(1)); + } +}; + +template +struct calculate_product> { + using type = ConstValue; + + template + KMM_HOST_DEVICE static type apply(const Shape& shape) { + return {}; + } +}; + +template> +struct ordered_strides_impl; + +template +struct ordered_strides_impl> { + using type = Strides>::type...>; + + KMM_HOST_DEVICE static type apply(Shape shape, IndexT alignment) { + if (N > 0 && alignment > 1) { + shape[0] = round_up_to_multiple(shape[0], alignment); + } + + return type(calculate_product>::apply(shape)...); + } +}; + +template +struct ordered_strides_impl> { + using type = + Strides>::type...>; + + KMM_HOST_DEVICE static type apply(Shape shape, IndexT alignment) { + if (N > 0 && alignment > 1) { + shape[N - 1] = round_up_to_multiple(shape[N - 1], alignment); + } + + return type(calculate_product>::apply(shape)...); + } +}; +} // namespace detail + +/// \addtogroup layout +/// @{ + +template +using make_strides_t = typename detail::ordered_strides_impl::type; + +/// Build a Strides object from a given shape and memory order (RowMajor or ColMajor). The +/// contiguous axis will always be `ConstValue`. +/// +/// For example, given a shape (10, 5, 3) +/// - RowMajor: strides are (IndexT(15), IndexT(3), ConstValue) +/// - ColMajor: strides are (ConstValue, IndexT(10), IndexT(50)) +template +KMM_HOST_DEVICE make_strides_t make_strides_from_shape( + const Shape& shape, + IndexT alignment = 1 +) { + return detail::ordered_strides_impl::apply(shape, alignment); +} + +/// @} + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + +namespace kmm { +template +std::ostream& operator<<(std::ostream& stream, const Strides& s) { + return stream << s.to_vec(); +} +} // namespace kmm + +template +struct fmt::formatter>: fmt::ostream_formatter {}; +#endif diff --git a/include/kmm/core/type_utils.hpp b/include/kmm/core/type_utils.hpp new file mode 100644 index 00000000..542c77cc --- /dev/null +++ b/include/kmm/core/type_utils.hpp @@ -0,0 +1,219 @@ +#pragma once + +#include "kmm/core/const_value.hpp" +#include "kmm/core/macros.hpp" + +namespace kmm { + +using size_t = decltype(sizeof(int)); +using ptrdiff_t = decltype(static_cast(nullptr) - static_cast(nullptr)); +using default_index_type = signed long long; // int64_t +using default_stride_type = ptrdiff_t; // + +namespace detail { +template +struct conditional_type_impl { + using type = TrueT; +}; + +template +struct conditional_type_impl { + using type = FalseT; +}; +} // namespace detail + +template +using conditional_t = typename detail::conditional_type_impl::type; + +namespace detail { +template +struct enable_if_type_impl {}; + +template +struct enable_if_type_impl { + using type = T; +}; +} // namespace detail + +template +using enable_if_t = typename detail::enable_if_type_impl::type; + +template +using assert_arity_t = typename detail::enable_if_type_impl::type; + +template +using identity_t = T; + +namespace detail { +// Usable only in unevaluated contexts (e.g. `decltype`); never defined/called. +template +T&& declval() noexcept; +} // namespace detail + +template +using void_t = void; + +namespace detail { +template +struct nonvoid_impl { + using type = T; +}; + +template +struct nonvoid_impl: nonvoid_impl {}; +} // namespace detail + +/// Resolves to the first non-`void` type in `Ts...`. +template +using nonvoid_t = typename detail::nonvoid_impl::type; + +template +struct TypeSequence {}; + +template +struct IndexSequence { + KMM_HOST_DEVICE + static constexpr size_t size() noexcept { + return sizeof...(Indices); + } + + /// Returns `fun(0) && fun(1) && fun(2) && ...` + template + KMM_HOST_DEVICE static bool all(F fun) { + return ((fun(ConstIndex())) && ...); + } + + /// Calls `fun(0); fun(1); fun(2); ...` + template + KMM_HOST_DEVICE static void for_each(F fun) { + ((fun(ConstIndex())), ...); + } + + /// Constructs `T` from `T{fun(0), fun(1), fun(2), ...}`. + template + KMM_HOST_DEVICE static T construct(F fun) { + return {fun(ConstIndex())...}; + } + + /// Constructs `T` from `T{value, value, value, ...}`. + template + KMM_HOST_DEVICE static T fill(const V& value) { + return {(ConstIndex(), value)...}; + } +}; + +namespace detail { +template +struct make_index_sequence_helper: make_index_sequence_helper {}; + +template +struct make_index_sequence_helper { + using type = IndexSequence; +}; + +// Shifts every index in `Seq` up by `Offset` (e.g. Offset=2, IndexSequence<0,1> -> IndexSequence<2,3>). +template +struct offset_index_sequence; + +template +struct offset_index_sequence> { + using type = IndexSequence<(Is + Offset)...>; +}; + +template> +struct reverse_index_sequence_helper { + using type = Dst; +}; + +template +struct reverse_index_sequence_helper, IndexSequence>: + reverse_index_sequence_helper, IndexSequence> {}; + +template +struct drop_index_sequence_helper; + +template +struct drop_index_sequence_helper> { + using type = IndexSequence<(Indices < Axis + 1 ? Indices - 1 : Indices)...>; +}; + +template +struct swap_index_sequence_helper; + +template +struct swap_index_sequence_helper> { + using type = IndexSequence<(Indices == I ? J : (Indices == J ? I : Indices))...>; +}; + +constexpr size_t move_axis_to_position_index(size_t p, size_t axis, size_t pos) { + // At position p, we get the new axis. + // All other axes either shift right (if p in axis...pos) or shift left (p in pos...axis). + return p == pos ? axis : p + size_t(axis <= p && p <= pos) - size_t(pos <= p && p <= axis); +} + +template +struct move_axis_to_position_index_sequence_helper; + +template +struct move_axis_to_position_index_sequence_helper> { + using type = IndexSequence; +}; +} // namespace detail + +/// Alias for IndexSequence<0, 1, ..., N-1> +template +using make_index_sequence = typename detail::make_index_sequence_helper<0, N>::type; + +/// Alias for IndexSequence +template +using range_index_sequence_t = typename detail::make_index_sequence_helper::type; + +/// Alias for IndexSequence +template +using reverse_index_sequence = + typename detail::reverse_index_sequence_helper>::type; + +/// Alias for IndexSequence<0, 1, ..., Axis-1, Axis+1, ..., N-1> +template +using drop_index_sequence = + typename detail::drop_index_sequence_helper>::type; + +/// Alias for IndexSequence<0, 1, ..., N-1> with the values at positions `I` and `J` swapped. +template +using swap_index_sequence = + typename detail::swap_index_sequence_helper>::type; + +/// Alias for IndexSequence<0, 1, ..., N-1> with axis `Axis` moved to position `Pos`, preserving +/// the relative order of the other axes. +template +using move_axis_to_position_index_sequence = typename detail:: + move_axis_to_position_index_sequence_helper>::type; + +namespace detail { +template +struct is_partial_permutation_impl {}; + +template +struct is_partial_permutation_impl, N> { + static constexpr bool value = true; +}; + +template +struct is_partial_permutation_impl, N> { + static constexpr bool value = ((First != Rest) && ...) && (First < N) + && is_partial_permutation_impl, N>::value; +}; +} // namespace detail + +/// True if `Seq` is an `IndexSequence` representing partial permutation for N. This means it +/// contains no duplicate indices and each index is less than N. +template +constexpr bool is_partial_permutation = detail::is_partial_permutation_impl::value; + +/// True if `Seq` is an `IndexSequence` representing a permutation of length N. This means it +/// contains N indices, no duplicates, and each index is less than N. In other words, it contains +/// the indices `0, 1, ..., N-1` in any arbitrary order. +template +constexpr bool is_permutation = is_partial_permutation; + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/core/vec.hpp b/include/kmm/core/vec.hpp new file mode 100644 index 00000000..8217575f --- /dev/null +++ b/include/kmm/core/vec.hpp @@ -0,0 +1,242 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/core/type_utils.hpp" + +namespace kmm { + +/// \addtogroup geometry +/// @{ + +/// A fixed-size array of `N` values of type `T`. +/// +/// This is the plain-storage building block used by `Point`, `Shape`, and `Bounds`: those +/// types inherit from `Vec` to add domain-specific semantics (coordinates, extents, ranges) +/// on top of simple indexed element access. Specializations exist for `N` in `{0, 1, 2, 3, 4}` +/// so that `x`/`y`/`z`/`w` accessors are available for low-dimensional vectors. +template +struct Vec { + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_ASSERT(i < N); + return values[i]; + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_ASSERT(i < N); + return values[i]; + } + + T values[N]; +}; + +template +struct Vec { + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_PANIC("index out of bounds"); + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_PANIC("index out of bounds"); + } +}; + +template +struct Vec { + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_DEBUG_ASSERT(i < 1); + return x; + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_DEBUG_ASSERT(i < 1); + return x; + } + + T x; +}; + +template +struct Vec { + KMM_HOST_DEVICE + constexpr Vec() {} + + KMM_HOST_DEVICE + constexpr Vec(const T& x, const T& y) : x(x), y(y) {} + + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_DEBUG_ASSERT(i < 2); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y}; + return this->*members[i]; + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_DEBUG_ASSERT(i < 2); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y}; + return this->*members[i]; + } + + T x; + T y; +}; + +template +struct Vec { + KMM_HOST_DEVICE + constexpr Vec() {} + + KMM_HOST_DEVICE + constexpr Vec(const T& x, const T& y, const T& z) : x(x), y(y), z(z) {} + + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_DEBUG_ASSERT(i < 3); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y, &Vec::z}; + return this->*members[i]; + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_DEBUG_ASSERT(i < 3); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y, &Vec::z}; + return this->*members[i]; + } + + T x; + T y; + T z; +}; + +template +struct Vec { + KMM_HOST_DEVICE + constexpr Vec() {} + + KMM_HOST_DEVICE + constexpr Vec(const T& x, const T& y, const T& z, const T& w) : x(x), y(y), z(z), w(w) {} + + KMM_HOST_DEVICE + constexpr T& operator[](size_t i) { + KMM_DEBUG_ASSERT(i < 4); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y, &Vec::z, &Vec::w}; + return this->*members[i]; + } + + KMM_HOST_DEVICE + constexpr const T& operator[](size_t i) const { + KMM_DEBUG_ASSERT(i < 4); + constexpr decltype(&Vec::x) members[] = {&Vec::x, &Vec::y, &Vec::z, &Vec::w}; + return this->*members[i]; + } + + T x; + T y; + T z; + T w; +}; + +template +Vec(const T&) -> Vec; + +template +Vec(const T&, const T&) -> Vec; + +template +Vec(const T&, const T&, const T&) -> Vec; + +template +Vec(const T&, const T&, const T&, const T&) -> Vec; + +template +KMM_HOST_DEVICE constexpr Vec fill(const T& value) { + return make_index_sequence::template fill>(value); +} + +/// @} + +namespace detail { +template +KMM_HOST_DEVICE Vec concat_impl( + const Vec& lhs, + const Vec& rhs, + IndexSequence, + IndexSequence +) { + return {lhs[Is]..., rhs[Js]...}; +} +} // namespace detail + +/// \addtogroup geometry +/// @{ + +template +KMM_HOST_DEVICE Vec concat(const Vec& lhs, const Vec& rhs) { + return detail::concat_impl(lhs, rhs, make_index_sequence(), make_index_sequence()); +} + +template +KMM_HOST_DEVICE bool operator==(const Vec& lhs, const Vec& rhs) { + bool result = true; + + if constexpr (N > 0) { + for (size_t i = 0; is_less(i, N); i++) { + result &= is_equal(lhs[i], rhs[i]); + } + } + + return result; +} + +template +KMM_HOST_DEVICE bool operator!=(const Vec& lhs, const Vec& rhs) { + return !(lhs == rhs); +} + +/// @} + +} // namespace kmm + +#if !KMM_IS_RTC + #include + + #include "fmt/ostream.h" + + #include "kmm/utils/hash_utils.hpp" + +namespace kmm { +template +std::ostream& operator<<(std::ostream& stream, const Vec& p) { + stream << "{"; + if constexpr (N > 0) { + stream << p[0]; + + for (size_t i = 1; is_less(i, N); i++) { + stream << ", " << p[i]; + } + } + return stream << "}"; +} +} // namespace kmm + +template +struct fmt::formatter>: fmt::ostream_formatter {}; + +template +struct std::hash> { + size_t operator()(const kmm::Vec& p) const { + size_t result = 0; + for (size_t i = 0; kmm::is_less(i, N); i++) { + kmm::hash_combine(result, p[i]); + } + return result; + } +}; +#endif \ No newline at end of file diff --git a/include/kmm/core/view.hpp b/include/kmm/core/view.hpp index 867e2e20..ac6f5713 100644 --- a/include/kmm/core/view.hpp +++ b/include/kmm/core/view.hpp @@ -1,923 +1,533 @@ #pragma once -#include "kmm/utils/fixed_vector.hpp" +#include "kmm/core/layout.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/core/type_utils.hpp" namespace kmm { -namespace views { -using default_index_type = signed long int; // int64_t -using default_stride_type = signed int; // int32_t - -template -struct static_domain { - static constexpr size_t rank = sizeof...(Dims); - using index_type = I; - - KMM_HOST_DEVICE - static static_domain from_domain(const static_domain& domain) noexcept { - return domain; - } - - KMM_HOST_DEVICE - constexpr index_type offset(size_t axis) const noexcept { - return static_cast(0); - } - - KMM_HOST_DEVICE - constexpr index_type size(size_t axis) const noexcept { - index_type sizes[rank + 1] = {Dims..., 0}; - return axis < rank ? sizes[axis] : static_cast(1); - } -}; - -template -struct static_offset { - static_assert(D::rank == sizeof...(Offsets), "Number of offsets must match rank of domain"); - - static constexpr size_t rank = D::rank; - using index_type = typename D::index_type; - - KMM_HOST_DEVICE - constexpr static_offset(D inner = {}) : m_inner(inner) {} - - template - KMM_HOST_DEVICE static_offset(const static_offset& domain) noexcept : - m_inner(domain.inner_domain()) {} - - KMM_HOST_DEVICE - constexpr index_type offset(size_t axis) const noexcept { - index_type offsets[rank + 1] = {Offsets..., 0}; - return m_inner.offset(axis) + (axis < rank ? offsets[axis] : static_cast(0)); - } - - KMM_HOST_DEVICE - constexpr index_type size(size_t axis) const noexcept { - return m_inner.size(axis); - } - - KMM_HOST_DEVICE - constexpr D inner_domain() const noexcept { - return m_inner; - } - - private: - D m_inner; -}; - -template -struct dynamic_domain { - static constexpr size_t rank = N; - using index_type = I; - - KMM_HOST_DEVICE - dynamic_domain(fixed_vector sizes = {}) noexcept : m_sizes(sizes) {} - - template - KMM_HOST_DEVICE dynamic_domain(static_domain) noexcept : - dynamic_domain(fixed_vector {Dims...}) {} - - KMM_HOST_DEVICE - constexpr index_type offset(size_t axis) const noexcept { - return static_cast(0); - } - - KMM_HOST_DEVICE - constexpr index_type size(size_t axis) const noexcept { - return axis < rank ? m_sizes[axis] : static_cast(1); - } - - private: - fixed_vector m_sizes; -}; - -template -struct dynamic_subdomain { - static constexpr size_t rank = N; - using index_type = I; - - KMM_HOST_DEVICE - constexpr dynamic_subdomain() noexcept { - for (size_t i = 0; i < rank; i++) { - m_sizes[i] = 0; - m_offsets[i] = 0; - } - } - - KMM_HOST_DEVICE - constexpr dynamic_subdomain(fixed_vector sizes) noexcept { - for (size_t i = 0; i < rank; i++) { - m_offsets[i] = 0; - m_sizes[i] = sizes[i]; - } - } - - KMM_HOST_DEVICE - constexpr dynamic_subdomain( - fixed_vector offsets, - fixed_vector sizes - ) noexcept { - for (size_t i = 0; i < rank; i++) { - m_offsets[i] = offsets[i]; - m_sizes[i] = sizes[i]; - } - } - - KMM_HOST_DEVICE - constexpr dynamic_subdomain(const dynamic_domain& domain) noexcept { - for (size_t i = 0; i < rank; i++) { - m_offsets[i] = 0; - m_sizes[i] = domain.size(i); - } - } - - template - KMM_HOST_DEVICE dynamic_subdomain(static_domain domain) noexcept : - dynamic_subdomain(dynamic_domain(domain)) {} - - template - KMM_HOST_DEVICE dynamic_subdomain(static_offset domain) noexcept : - dynamic_subdomain(domain.inner_domain()) { - for (size_t i = 0; i < rank; i++) { - m_offsets[i] = static_cast(domain.offset(i)); - m_sizes[i] = static_cast(domain.size(i)); - } - } - - KMM_HOST_DEVICE - constexpr index_type offset(size_t axis) const noexcept { - return axis < rank ? m_offsets[axis] : static_cast(0); - } - - KMM_HOST_DEVICE - constexpr index_type size(size_t axis) const noexcept { - return axis < rank ? m_sizes[axis] : static_cast(1); +/// Accessor tag: no restriction -- the view's data may be dereferenced from both host and device +/// code. This is the default. +struct AnyAccessor { + template + KMM_HOST_DEVICE T& dereference(T* input) const noexcept { + return *input; } - - private: - fixed_vector m_offsets; - fixed_vector m_sizes; }; -template -struct static_mapping { - static constexpr size_t rank = sizeof...(Strides); - using stride_type = S; - - KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - S strides[rank] = {Strides...}; - return axis < rank ? strides[axis] : static_cast(0); +/// Accessor tag: the view's data may only be dereferenced from device (GPU) code. Dereferencing +/// it from host code panics at runtime with a clear message. +struct DeviceAccessor { + template + KMM_HOST_DEVICE T& dereference(T* input) const noexcept { +#if KMM_IS_DEVICE + return *input; +#else + KMM_PANIC("cannot access device data from host code"); +#endif } }; -template -struct static_mapping { - static constexpr size_t rank = 0; - using stride_type = S; +namespace detail { - KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - return static_cast(0); - } -}; +// Tag for NDView's raw (pointer, layout) constructor: pointer is already offset-adjusted, so +// it must NOT be shifted again by layout.base_offset(). A free type (not a per-NDView +// nested type) so it's the same type across every NDView instantiation -- +// with_layout() below constructs a NDView that differs from the enclosing +// NDView. +struct view_raw_ctor_tag {}; -template -using linear_mapping = static_mapping(1)>; - -template -struct contiguous_axis_mapping { - static_assert(ContAxis < N, "Axis cannot exceed dimensionality"); - static constexpr size_t rank = N; - using stride_type = S; +template +class ViewAccessor { + public: + static constexpr size_t rank = ViewT::rank; + using index_type = typename ViewT::index_type; + using ndindex_type = typename ViewT::ndindex_type; KMM_HOST_DEVICE - explicit constexpr contiguous_axis_mapping(fixed_vector strides) noexcept { - for (size_t i = 0; i < rank - 1; i++) { - m_strides[i] = i < ContAxis ? strides[i] : strides[i + 1]; - } + constexpr ViewAccessor(const ViewT* view, identity_t... p) : m_view(view) { + ((m_point[Is] = p), ...); } KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - if (axis < ContAxis) { - return m_strides[axis]; - } else if (axis == ContAxis) { - return static_cast(1); - } else if (axis < rank) { - return m_strides[axis - 1]; + static decltype(auto) index(const ViewT* view, identity_t... p) noexcept { + if constexpr (sizeof...(Is) == rank) { + return (*view)(p...); } else { - return static_cast(0); + return ViewAccessor(view, p...); } } - private: - stride_type m_strides[rank - 1]; -}; - -template -struct contiguous_axis_mapping<1, 0, S> { - static constexpr size_t rank = 1; - using stride_type = S; - - KMM_HOST_DEVICE - explicit constexpr contiguous_axis_mapping(fixed_vector strides) noexcept {} - - KMM_HOST_DEVICE constexpr contiguous_axis_mapping(static_mapping m) noexcept {} - - KMM_HOST_DEVICE - operator static_mapping() { - return {}; - } - KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - return axis == 0 ? static_cast(1) : static_cast(0); - } -}; - -template -struct strided_mapping { - static constexpr size_t rank = N; - using stride_type = S; - - template - KMM_HOST_DEVICE static constexpr strided_mapping from_mapping(const M& mapping) noexcept { - fixed_vector strides; - - for (size_t i = 0; i < N; i++) { - strides[i] = static_cast(mapping.stride(i)); - } - - return {strides}; - } - - KMM_HOST_DEVICE - constexpr strided_mapping() noexcept { - for (size_t i = 0; i < N; i++) { - m_strides[i] = static_cast(0); - } - } - - KMM_HOST_DEVICE - constexpr strided_mapping(fixed_vector strides) noexcept { - for (size_t i = 0; i < N; i++) { - m_strides[i] = strides[i]; - } - } - - template - KMM_HOST_DEVICE constexpr strided_mapping(contiguous_axis_mapping m - ) noexcept : - strided_mapping(from_mapping(m)) {} - - template - KMM_HOST_DEVICE constexpr strided_mapping(static_mapping m) noexcept : - strided_mapping(from_mapping(m)) {} - - KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - return axis < rank ? m_strides[axis] : static_cast(0); + decltype(auto) operator[](index_type i) const noexcept { + static constexpr size_t Axis = sizeof...(Is); + return ViewAccessor::index(m_view, m_point[Is]..., i); } private: - fixed_vector m_strides; + const ViewT* m_view; + index_type m_point[rank] = {}; }; -template -struct strided_mapping<0, S> { - static constexpr size_t rank = 0; - using stride_type = S; - - KMM_HOST_DEVICE - constexpr strided_mapping(fixed_vector strides = {}) noexcept {} - - KMM_HOST_DEVICE constexpr strided_mapping(static_mapping m) noexcept {} - - KMM_HOST_DEVICE - operator static_mapping() { - return {}; - } - - KMM_HOST_DEVICE - stride_type stride(size_t axis) const noexcept { - return static_cast(0); +template +class NDViewBase { + public: + // The `typename = ...` constrains this to index types that convert directly to a scalar + // `index_type` (e.g. `int`), so it does not compete with `NDView`'s non-template + // `operator[](const ndindex_type&)` overload for `Vec`/`Point` indices (which do not convert + // to a scalar `index_type`). `Self` (defaulted to `DerivedT`) makes the default argument + // depend on this function template's own parameters rather than solely on the enclosing + // class template's `DerivedT`, deferring its instantiation to call time -- `DerivedT` is + // still incomplete while `NDViewBase` is being instantiated as its base class. + template< + typename IndexT, + typename Self = DerivedT, + typename = decltype(static_cast(declval()))> + KMM_HOST_DEVICE decltype(auto) operator[](IndexT i) const noexcept { + const auto* self = static_cast(this); + using index_type = typename DerivedT::index_type; + return ViewAccessor {self}[static_cast(i)]; } }; -namespace details { -template -struct select_contiguous_axis_mapping { - using type = contiguous_axis_mapping; -}; - -template -struct select_contiguous_axis_mapping<0, ConstAxis, S> { - using type = strided_mapping<0, S>; +// Rank 0 has no axis to chain-index into. Deleted (rather than simply absent) so that +// NDView's `using base_type::operator[];` always has a member to name, regardless +// of rank; actually calling this is a compile error with a clear cause. +template +class NDViewBase { + public: + template + void operator[](IndexT) const = delete; }; -} // namespace details - -template -struct left_to_right_layout { - template - using mapping_type = typename details::select_contiguous_axis_mapping::type; - - template - KMM_HOST_DEVICE static mapping_type from_domain(const D& domain) noexcept { - fixed_vector strides; - S stride = 1; - - for (size_t i = 0; i < D::rank; i++) { - strides[i] = stride; - stride *= static_cast(domain.size(i)); - } - return mapping_type(strides); - } -}; +} // namespace detail -template -struct right_to_left_layout { - template - using mapping_type = - typename details::select_contiguous_axis_mapping::type; +/// \addtogroup views +/// @{ - template - KMM_HOST_DEVICE static mapping_type from_domain(const D& domain) noexcept { - fixed_vector strides; - S stride = 1; +/// A dense or strided view over a multi-dimensional array. +template +class NDView: public detail::NDViewBase, LayoutT::rank> { + using base_type = detail::NDViewBase, LayoutT::rank>; - for (size_t i = D::rank; i > 0; i--) { - strides[i - 1] = stride; - stride *= static_cast(domain.size(i - 1)); - } + public: + using base_type::operator[]; - return mapping_type(strides); - } -}; + using self_type = NDView; + using layout_type = LayoutT; + using element_type = T; + using pointer = T*; + using reference = T&; + using accessor_type = AccessorT; -template -struct strided_layout { - template - using mapping_type = strided_mapping; + static constexpr size_t rank = layout_type::rank; + using index_type = typename layout_type::index_type; + using ndindex_type = typename layout_type::ndindex_type; + using shape_type = typename layout_type::shape_type; + using bounds_type = typename layout_type::bounds_type; + using stride_type = typename layout_type::stride_type; + using ndstrides_type = typename layout_type::ndstrides_type; - template - KMM_HOST_DEVICE static mapping_type from_domain(const D& domain) noexcept { - return mapping_type::from_mapping(right_to_left_layout::from_domain(domain)); - } -}; - -template -struct static_layout { - template - using mapping_type = static_mapping; + template + using rebind_layout = NDView; - template - KMM_HOST_DEVICE static mapping_type from_domain(const D& domain) noexcept { - static_assert(sizeof...(Strides) == D::rank, "number of strides must match dimensionality"); - return {}; - } -}; + template + using drop_axis_type = rebind_layout>; -using default_layout = right_to_left_layout<>; + template + using insert_axis_type = rebind_layout>; -template -struct drop_axis_domain { - static_assert(Axis < D::rank); - using index_type = typename D::index_type; - using type = dynamic_subdomain; + using reverse_axes_type = rebind_layout; - KMM_HOST_DEVICE - static type call(const D& domain) noexcept { - fixed_vector new_offsets; - fixed_vector new_sizes; - size_t axis = Axis; - - for (size_t i = 0; i < D::rank - 1; i++) { - new_offsets[i] = domain.offset(i < axis ? i : i + 1); - new_sizes[i] = domain.size(i < axis ? i : i + 1); - } + template + using permute_axes_type = + rebind_layout>; - return {new_offsets, new_sizes}; - } -}; + template + using swap_axes_type = rebind_layout>; -template -struct drop_axis_domain> { - static_assert(DropAxis < N); - using index_type = I; - using type = dynamic_domain; + using transpose_type = rebind_layout; - KMM_HOST_DEVICE - static type call(const dynamic_domain& domain) noexcept { - fixed_vector new_sizes; + template + using move_axis_to_position_type = + rebind_layout>; - for (size_t i = 0; i < N - 1; i++) { - new_sizes[i] = domain.size(i < DropAxis ? i : i + 1); - } + template + using move_axis_to_front_type = + rebind_layout>; - return {new_sizes}; - } -}; + template + using move_axis_to_back_type = + rebind_layout>; -template -struct drop_axis_domain>: - drop_axis_domain> {}; + using zero_origin_type = rebind_layout; + using move_origin_type = rebind_layout; -template -struct drop_axis_layout { - using old_mapping_type = typename L::template mapping_type; - using stride_type = typename old_mapping_type::stride_type; + template + using slice_axis_type = + rebind_layout>; - using new_mapping_type = strided_mapping; - using type = strided_layout; + template + using slice_type = rebind_layout>; + /// Constructs an empty (null) view. KMM_HOST_DEVICE - static new_mapping_type call(const old_mapping_type& mapping) noexcept { - fixed_vector new_strides; - - for (size_t i = 0; i < D::rank - 1; i++) { - new_strides[i] = mapping.stride(i < DropAxis ? i : i + 1); - } - - return {new_strides}; - } -}; - -template -struct drop_axis_layout<0, right_to_left_layout, D> { - using old_mapping_type = contiguous_axis_mapping; - using stride_type = S; - - using new_mapping_type = - typename details::select_contiguous_axis_mapping::type; - using type = right_to_left_layout; + constexpr NDView() = default; + /// Constructs a view from a pointer already adjusted for the layout's base offset. KMM_HOST_DEVICE - static new_mapping_type call(const old_mapping_type& mapping) noexcept { - fixed_vector new_strides; - - for (size_t i = 0; i < D::rank - 1; i++) { - new_strides[i] = mapping.stride(i + 1); - } - - return new_mapping_type {new_strides}; - } -}; - -struct host_accessor { - template - KMM_HOST_DEVICE T& dereference_pointer(T* ptr) const noexcept { - return *ptr; - } -}; - -struct device_accessor { - template - KMM_HOST_DEVICE T& dereference_pointer(T* ptr) const { -#if __CUDA_ARCH__ or __HIP_DEVICE_COMPILE__ - return *ptr; -#else - throw std::runtime_error("device data cannot be accessed on host"); -#endif - } -}; - -template -struct convert_pointer; - -template -struct convert_pointer { - static KMM_HOST_DEVICE T* call(T* p) { - return p; - } -}; - -template -struct convert_pointer: convert_pointer {}; - -} // namespace views -template -struct ViewSubscript { - using type = ViewSubscript; - using subscript_type = typename ViewSubscript::type; - using index_type = typename D::index_type; - using ndindex_type = fixed_vector; + constexpr NDView( + detail::view_raw_ctor_tag, + pointer data, + layout_type layout, + AccessorT accessor = {} + ) : + m_data(data), + m_layout(layout), + m_accessor(accessor) {} + /// Constructs a view over the given data pointer and layout. KMM_HOST_DEVICE - static type instantiate(const View* base, ndindex_type index = {}) noexcept { - return type {base, index}; - } + constexpr NDView(pointer data, layout_type layout, AccessorT accessor = {}) : + NDView(detail::view_raw_ctor_tag {}, data + layout.base_offset(), layout, accessor) {} - KMM_HOST_DEVICE - ViewSubscript(const View* base, ndindex_type index) noexcept : base_(base), index_(index) {} + /// Converting constructor from a view over a compatible element/layout type with the same + /// accessor tag (converting between AnyAccessor and DeviceAccessor is not allowed). Carries + /// over the source view's accessor value rather than default-constructing a new one, so any + /// accessor state survives the conversion. + template + KMM_HOST_DEVICE constexpr NDView(const NDView& that) : + NDView(that.data(), that.layout(), that.accessor()) {} + /// Returns the underlying data pointer. KMM_HOST_DEVICE - subscript_type operator[](index_type index) { - index_[K] = index; - return ViewSubscript::instantiate(base_, index_); + constexpr pointer data() const noexcept { + return m_data - m_layout.base_offset(); } - private: - const View* base_; - ndindex_type index_; -}; - -template -struct ViewSubscript { - using type = T&; - using index_type = typename D::index_type; - using ndindex_type = fixed_vector; - + /// Returns the underlying data pointer at the given index. KMM_HOST_DEVICE - static type instantiate(const View* base, ndindex_type index) { - return base->access(index); + constexpr pointer data_at(ndindex_type index) const noexcept { + return &m_data[m_layout.local_offset(index)]; } -}; - -template -struct AbstractViewBase { - using index_type = typename D::index_type; - using subscript_type = typename ViewSubscript::subscript_type; + /// Returns the layout describing this view's domain, strides, and base offset. KMM_HOST_DEVICE - subscript_type operator[](index_type index) const { - return ViewSubscript::instantiate(static_cast(this))[index]; + constexpr const layout_type& layout() const noexcept { + return m_layout; } -}; - -template -struct AbstractViewBase { - using reference = T&; + /// Returns the accessor used to dereference this view's data. KMM_HOST_DEVICE - reference operator*() const { - return static_cast(this)->access({}); + constexpr accessor_type accessor() const noexcept { + return m_accessor; } -}; - -template -struct AbstractView: - public D, - public L::template mapping_type, - public A, - public AbstractViewBase, T, D> { - using self_type = AbstractView; - using value_type = T; - using domain_type = D; - using layout_type = L; - using mapping_type = typename L::template mapping_type; - using accessor_type = A; - using pointer = T*; - using reference = T&; - - static constexpr size_t rank = D::rank; - using index_type = typename domain_type::index_type; - using stride_type = typename mapping_type::stride_type; - using ndindex_type = fixed_vector; - using ndstride_type = fixed_vector; - - using origin_domain_type = views::dynamic_domain; - using shifted_domain_type = views::dynamic_subdomain; - - AbstractView(const AbstractView&) = default; - AbstractView(AbstractView&&) noexcept = default; - - AbstractView& operator=(const AbstractView&) = default; - AbstractView& operator=(AbstractView&&) noexcept = default; + /// Returns the extent (size) along each axis. KMM_HOST_DEVICE - AbstractView( - pointer data, - domain_type domain, - mapping_type mapping, - accessor_type accessor = {} - ) noexcept : - domain_type(domain), - mapping_type(mapping), - accessor_type(accessor) { - m_data = data - this->linearize_index(offsets()); + shape_type shape() const noexcept { + return m_layout.shape(); } + /// Returns the bounds (begin/end per axis) covered by this view. KMM_HOST_DEVICE - AbstractView(pointer data = nullptr, domain_type domain = {}) noexcept : - AbstractView(data, domain, layout_type::from_domain(domain)) {} - - template< - typename T2, - typename D2, - typename L2, - typename = decltype(views::convert_pointer::call(nullptr))> - KMM_HOST_DEVICE AbstractView(const AbstractView& that) noexcept : - AbstractView( - views::convert_pointer::call(that.data()), - domain_type(that.domain()), - mapping_type(that.mapping()), - that.accessor() - ) {} - - template - KMM_HOST_DEVICE AbstractView& operator=(const AbstractView& that) noexcept { - return *this = AbstractView(that); + bounds_type bounds() const noexcept { + return m_layout.bounds(); } + /// Returns the extent along the given axis. KMM_HOST_DEVICE - pointer data() const noexcept { - return data_at(offsets()); + index_type extent(size_t axis) const noexcept { + return m_layout.extent(axis); } + /// Returns the total number of elements covered by this view. KMM_HOST_DEVICE - operator pointer() const noexcept { - return data(); + index_type size() const noexcept { + return m_layout.size(); } + /// Returns whether this view covers zero elements. KMM_HOST_DEVICE - const mapping_type& mapping() const noexcept { - return *this; + bool is_empty() const noexcept { + return m_layout.is_empty(); } + /// Returns whether the given index falls within this view's bounds. KMM_HOST_DEVICE - const domain_type& domain() const noexcept { - return *this; + bool contains(const ndindex_type& index) const noexcept { + return m_layout.contains(index); } + /// Returns the stride along the given axis. KMM_HOST_DEVICE - const accessor_type& accessor() const noexcept { - return *this; + stride_type stride(size_t axis) const noexcept { + return m_layout.stride(axis); } + /// Returns the stride along each axis. KMM_HOST_DEVICE - index_type size(size_t axis) const noexcept { - return domain().size(axis); + ndstrides_type strides() const noexcept { + return m_layout.strides(); } + /// Returns whether the view's strides are contiguous in the given memory order. KMM_HOST_DEVICE - index_type size() const noexcept { - index_type volume = 1; - for (size_t i = 0; i < rank; i++) { - volume *= domain().size(i); - } - return volume; + bool is_contiguous(MemoryOrder order = MemoryOrder::RowMajor) const noexcept { + return m_layout.is_contiguous(order); } + /// Returns a reference to the element at the given index, asserting it is in bounds. KMM_HOST_DEVICE - size_t size_in_bytes() const noexcept { - size_t nbytes = sizeof(T); - for (size_t i = 0; i < rank; i++) { - nbytes *= static_cast(domain().size(i)); - } - return nbytes; + reference access(const ndindex_type& index) const noexcept { + KMM_DEBUG_ASSERT(contains(index)); + return m_accessor.dereference(data_at(index)); } - KMM_HOST_DEVICE - stride_type stride(size_t axis = 0) const noexcept { - return mapping().stride(axis); + KMM_HOST_DEVICE reference operator[](const ndindex_type& index) const noexcept { + return access(index); } - KMM_HOST_DEVICE - index_type offset(size_t axis = 0) const noexcept { - return domain().offset(axis); + /// Returns a reference to the element at the given per-axis indices. + template> + KMM_HOST_DEVICE reference operator()(Indices... indices) const noexcept { + return access(ndindex_type {static_cast(indices)...}); } + /// Returns this view rebased so its domain starts at the zero index. KMM_HOST_DEVICE - ndstride_type strides() const noexcept { - ndstride_type result; - for (size_t i = 0; i < rank; i++) { - result[i] = stride(i); - } - return result; + zero_origin_type zero_origin() const noexcept { + return with_layout(m_layout.zero_origin()); } + /// Returns this view shifted so it originates at the given index, keeping the same shape. KMM_HOST_DEVICE - ndindex_type offsets() const noexcept { - ndindex_type result; - for (size_t i = 0; i < rank; i++) { - result[i] = offset(i); - } - return result; + move_origin_type move_origin(ndindex_type new_origin) const noexcept { + return with_layout(m_layout.move_origin(new_origin)); } + /// Returns this view restricted to the intersection of its bounds and the given bounds. KMM_HOST_DEVICE - ndindex_type sizes() const noexcept { - ndindex_type result; - for (size_t i = 0; i < rank; i++) { - result[i] = this->size(i); - } - return result; + move_origin_type restrict_bounds(bounds_type new_bounds) const noexcept { + return with_layout(m_layout.restrict_bounds(new_bounds)); } - KMM_HOST_DEVICE - index_type begin(size_t axis = 0) const noexcept { - return offset(axis); + /// Returns this view restricted along one axis to the intersection with [start, stop). + template + KMM_HOST_DEVICE move_origin_type + restrict_axis(index_type start, index_type stop) const noexcept { + return with_layout(m_layout.template restrict_axis(start, stop)); } - KMM_HOST_DEVICE - index_type end(size_t axis = 0) const noexcept { - return begin(axis) + this->size(axis); + /// Returns this view with the given axis dropped, fixed at the given index. + template + KMM_HOST_DEVICE drop_axis_type drop_axis(index_type index) const noexcept { + return with_layout(m_layout.template drop_axis(index)); } - template - KMM_HOST_DEVICE P linearize_index(ndindex_type ndindex, P base = {}) const noexcept { - for (size_t i = 0; i < rank; i++) { - base += - static_cast(ndindex[i]) * static_cast(mapping().stride(i)); - } - - return base; + /// Returns this view with a new broadcast axis of the given extent inserted at the given position. + template + KMM_HOST_DEVICE insert_axis_type insert_axis( + index_type extent = static_cast(1) + ) const noexcept { + return with_layout(m_layout.template insert_axis(extent)); } - KMM_HOST_DEVICE - value_type* data_at(ndindex_type ndindex) const noexcept { - return linearize_index(ndindex, m_data); + /// Returns this view with the order of all axes reversed. + KMM_HOST_DEVICE reverse_axes_type reverse_axes() const noexcept { + return with_layout(m_layout.reverse_axes()); } - template - KMM_HOST_DEVICE value_type* data_at(Indices... indices) const noexcept { - static_assert(sizeof...(Indices) == rank, "invalid number of indices"); - return data_at(ndindex_type {indices...}); + /// Returns this view with its axes reordered according to the given permutation, e.g. + /// `permute_axes<2, 0, 1>()` moves the current axis 2 to position 0, axis 0 to position 1, + /// and axis 1 to position 2. + template + KMM_HOST_DEVICE permute_axes_type permute_axes( + IndexSequence seq = {} + ) const noexcept { + return with_layout(m_layout.template permute_axes(seq)); } - KMM_HOST_DEVICE - reference access(ndindex_type ndindex) const noexcept { - return accessor().dereference_pointer(data_at(ndindex)); + /// Returns this view with axes `I` and `J` swapped. + template + KMM_HOST_DEVICE swap_axes_type swap_axes() const noexcept { + return with_layout(m_layout.template swap_axes()); } - template - KMM_HOST_DEVICE reference operator()(Indices... indices) const noexcept { - static_assert(sizeof...(Indices) == rank, "invalid number of indices"); - return access(ndindex_type {indices...}); + /// Returns this view with axes 0 and 1 swapped. Only valid for a rank-2 view; use + /// `swap_axes` or `permute_axes` for other ranks. + KMM_HOST_DEVICE transpose_type transpose() const noexcept { + return with_layout(m_layout.transpose()); } - KMM_HOST_DEVICE - bool is_empty() const noexcept { - bool result = false; - for (size_t i = 0; i < rank; i++) { - result |= domain().size(i) <= static_cast(0); - } - return result; + /// Returns this view with the given axis moved to the given position, preserving the + /// relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_position_type move_axis_to_position() const noexcept { + return with_layout(m_layout.template move_axis_to_position()); } - KMM_HOST_DEVICE - bool in_bounds(ndindex_type ndindex) const noexcept { - bool result = true; - for (size_t i = 0; i < rank; i++) { - result &= ndindex[i] >= domain().offset(i); - result &= ndindex[i] - domain().offset(i) < domain().size(i); - } - return result; + /// Returns this view with the given axis moved to the front (position 0), preserving the + /// relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_front_type move_axis_to_front() const noexcept { + return with_layout(m_layout.template move_axis_to_front()); } - template - KMM_HOST_DEVICE bool in_bounds(Indices... indices) const noexcept { - static_assert(sizeof...(Indices) == rank, "invalid number of indices"); - return in_bounds(ndindex_type {indices...}); + /// Returns this view with the given axis moved to the back (position `rank - 1`), + /// preserving the relative order of the remaining axes. + template + KMM_HOST_DEVICE move_axis_to_back_type move_axis_to_back() const noexcept { + return with_layout(m_layout.template move_axis_to_back()); } - KMM_HOST_DEVICE - bool is_contiguous() const noexcept { - stride_type curr = 1; - bool result = true; - - for (size_t i = 0; i < rank; i++) { - result &= mapping().stride(rank - i - 1) == curr; - curr *= static_cast(domain().size(rank - i - 1)); - } - - return result; - } - - template - KMM_HOST_DEVICE AbstractView< - value_type, - typename views::drop_axis_domain::type, - typename views::drop_axis_layout::type, - accessor_type> - drop_axis(index_type index) const noexcept { - static_assert(Axis < rank, "axis out of bounds"); - return AbstractView< - value_type, - typename views::drop_axis_domain::type, - typename views::drop_axis_layout::type, - accessor_type> { - data() - mapping().stride(Axis) * offset(Axis) + mapping().stride(Axis) * index, - views::drop_axis_domain::call(domain()), - views::drop_axis_layout::call(mapping()), - accessor() - }; - } - - template - KMM_HOST_DEVICE AbstractView< - value_type, - typename views::drop_axis_domain::type, - typename views::drop_axis_layout::type, - accessor_type> - drop_axis() const noexcept { - static_assert(Axis < rank, "axis out of bounds"); - return this->template drop_axis(offset(Axis)); - } - - AbstractView // - shift_to_origin() const noexcept { - auto new_domain = views::dynamic_domain(sizes()); - return {data(), new_domain, mapping(), accessor()}; - } - - AbstractView // - shift_to(ndindex_type new_offsets) const noexcept { - auto new_domain = views::dynamic_subdomain(new_offsets, sizes()); - return {data(), new_domain, mapping(), accessor()}; - } - - AbstractView // - shift_by(ndindex_type amount) const noexcept { - auto new_offsets = offsets(); - for (size_t i = 0; i < rank; i++) { - new_offsets[i] += amount[i]; - } - return shift_to(new_offsets); + /// Returns this view with the given axis sliced according to the given slice token (e.g. `all`, a `Range`, `new_axis`). + template + slice_axis_type slice_axis(const SliceT& slice) const noexcept { + return with_layout(m_layout.template slice_axis(slice)); } + /// Returns this view with the given axis narrowed to the range [start, end). template - AbstractView // - shift_axis_to(index_type new_offset) const noexcept { - static_assert(Axis < rank, "axis out of bounds"); - auto new_offsets = offsets(); - new_offsets[Axis] = new_offset; - return shift_to(new_offsets); + KMM_HOST_DEVICE self_type slice_axis(index_type start, index_type end) const noexcept { + return with_layout(m_layout.template slice_axis(start, end)); } - template - AbstractView // - shift_axis_by(index_type amount) const noexcept { - return shift_axis_to(offset(Axis) + amount); + /// Returns this view sliced across all axes at once, one slice token per axis. + template + slice_type slice(const Slices&... slices) const noexcept { + return with_layout(m_layout.slice(slices...)); } private: - pointer m_data; + template + rebind_layout with_layout(NewLayoutT&& new_layout) const noexcept { + auto delta = new_layout.base_offset() - m_layout.base_offset(); + return {detail::view_raw_ctor_tag {}, m_data + delta, new_layout, m_accessor}; + } + + pointer m_data = nullptr; + KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS LayoutT m_layout {}; + KMM_ATTRIBUTE_NO_UNIQUE_ADDRESS AccessorT m_accessor {}; }; +/// A read-only view over a Shape domain with the given policy, for the common case +/// where callers just want an N-dimensional dense/strided view whose domain starts at index 0. +/// Use `ViewMut` for a mutable view. template< typename T, size_t N = 1, - typename L = views::default_layout, - typename A = views::host_accessor> -using View = AbstractView, L, A>; + typename PolicyT = RowMajor, + typename IndexT = default_index_type, + typename AccessorT = AnyAccessor> +using View = NDView, PolicyT>, AccessorT>; +/// A mutable view over a Shape domain with the given policy. See `View` for the +/// read-only counterpart. template< typename T, size_t N = 1, - typename L = views::default_layout, - typename A = views::host_accessor> -using ViewMut = AbstractView, L, A>; - + typename PolicyT = RowMajor, + typename IndexT = default_index_type, + typename AccessorT = AnyAccessor> +using ViewMut = NDView, PolicyT>, AccessorT>; + +/// A read-only view over a Bounds domain with the given policy -- a view over a +/// sub-region/window of a larger domain, whose origin need not start at index 0 (unlike `View`, +/// which always starts at index 0). Use `SubViewMut` for a mutable view. template< typename T, size_t N = 1, - typename L = views::default_layout, - typename A = views::host_accessor> -using Subview = AbstractView, L, A>; + typename PolicyT = RowMajor, + typename IndexT = default_index_type, + typename AccessorT = AnyAccessor> +using SubView = NDView, PolicyT>, AccessorT>; +/// A mutable view over a Bounds domain with the given policy. See `SubView` for the +/// read-only counterpart. template< typename T, size_t N = 1, - typename L = views::default_layout, - typename A = views::host_accessor> -using SubviewMut = AbstractView, L, A>; - -template -using ViewStrided = View, A>; - -template -using ViewStridedMut = ViewMut, A>; - -template -using SubviewStrided = Subview, A>; - -template -using SubviewStridedMut = SubviewMut, A>; - -template -using GPUView = View; + typename PolicyT = RowMajor, + typename IndexT = default_index_type, + typename AccessorT = AnyAccessor> +using SubViewMut = NDView, PolicyT>, AccessorT>; + +/// A read-only view over a Shape domain whose data may only be dereferenced from +/// device (GPU) code. Use `DeviceViewMut` for a mutable view, or `DeviceSubView` for the +/// arbitrary-origin (Bounds) counterpart. +template< + typename T, + size_t N = 1, + typename PolicyT = RowMajor, + typename IndexT = default_index_type> +using DeviceView = View; -template -using GPUViewMut = ViewMut; +/// A mutable view over a Shape domain whose data may only be dereferenced from device +/// (GPU) code. +template< + typename T, + size_t N = 1, + typename PolicyT = RowMajor, + typename IndexT = default_index_type> +using DeviceViewMut = ViewMut; -template -using GPUViewStrided = ViewStrided; +/// A read-only view over a Bounds domain whose data may only be dereferenced from +/// device (GPU) code. Use `DeviceSubViewMut` for a mutable view. +template< + typename T, + size_t N = 1, + typename PolicyT = RowMajor, + typename IndexT = default_index_type> +using DeviceSubView = SubView; -template -using GPUViewStridedMut = ViewStridedMut; +/// A mutable view over a Bounds domain whose data may only be dereferenced from device +/// (GPU) code. +template< + typename T, + size_t N = 1, + typename PolicyT = RowMajor, + typename IndexT = default_index_type> +using DeviceSubViewMut = SubViewMut; -template -using GPUSubview = Subview; +/// Constructs a mutable ViewMut over the given data pointer and shape. Whether the result is +/// read-only follows T's own constness (e.g. passing a `const int*` yields a read-only view), +/// exactly like ViewMut does. +template< + typename PolicyT = RowMajor, + typename AccessorT = AnyAccessor, + typename T, + size_t N, + typename IndexT = default_index_type> +KMM_HOST_DEVICE ViewMut make_view( + T* data, + Shape shape, + PolicyT policy = {} +) { + return {data, make_layout(shape, policy)}; +} + +/// @} -template -using GPUSubviewMut = SubviewMut; +} // namespace kmm -template -using GPUSubviewStrided = SubviewStrided; +// NDView has no base class to delegate to (other than the internal ViewAccessor/NDViewBase +// storage helpers), so this prints/hashes it as the tuple of its two constituent parts (the raw +// pointer and the layout), matching how Layout itself is printed. +#if !KMM_IS_RTC + #include -template -using GPUSubviewStridedMut = SubviewStridedMut; + #include "fmt/ostream.h" +namespace kmm { +template +std::ostream& operator<<(std::ostream& stream, const NDView& view) { + return stream << "NDView(data=" << static_cast(view.data()) + << ", layout=" << view.layout() << ")"; +} } // namespace kmm + +template +struct fmt::formatter>: fmt::ostream_formatter {}; +#endif diff --git a/include/kmm/kmm.hpp b/include/kmm/kmm.hpp index b8081d54..d10a82cb 100644 --- a/include/kmm/kmm.hpp +++ b/include/kmm/kmm.hpp @@ -1,11 +1,2 @@ -#include "kmm/api/access.hpp" -#include "kmm/api/argument.hpp" -#include "kmm/api/array.hpp" -#include "kmm/api/launcher.hpp" -#include "kmm/api/mapper.hpp" -#include "kmm/api/parallel_submit.hpp" -#include "kmm/api/runtime_handle.hpp" -#include "kmm/api/task_group.hpp" -#include "kmm/api/view_argument.hpp" -#include "kmm/core/distribution.hpp" -#include "kmm/core/domain.hpp" +#include "kmm/api/dist_array.hpp" +#include "kmm/api/host.hpp" diff --git a/include/kmm/memops/gpu_copy.hpp b/include/kmm/memops/gpu_copy.hpp deleted file mode 100644 index 35867868..00000000 --- a/include/kmm/memops/gpu_copy.hpp +++ /dev/null @@ -1,43 +0,0 @@ -#pragma once - -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -void execute_gpu_h2d_copy_async( - g_stream_t stream, - const void* src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -); - -void execute_gpu_d2h_copy_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - void* dst_buffer, - CopyDef copy_description -); - -void execute_gpu_d2d_copy_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -); - -void execute_gpu_h2d_copy( - const void* src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -); - -void execute_gpu_d2h_copy(g_device_ptr_t src_buffer, void* dst_buffer, CopyDef copy_description); - -void execute_gpu_d2d_copy( - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -); - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/memops/gpu_fill.hpp b/include/kmm/memops/gpu_fill.hpp deleted file mode 100644 index 69cd692c..00000000 --- a/include/kmm/memops/gpu_fill.hpp +++ /dev/null @@ -1,10 +0,0 @@ -#pragma once - -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -void execute_gpu_fill_async(g_stream_t stream, g_device_ptr_t dst_buffer, const FillDef& fill); - -} \ No newline at end of file diff --git a/include/kmm/memops/gpu_reduction.hpp b/include/kmm/memops/gpu_reduction.hpp deleted file mode 100644 index 15f46bf4..00000000 --- a/include/kmm/memops/gpu_reduction.hpp +++ /dev/null @@ -1,19 +0,0 @@ -#pragma once - -#include "kmm/core/backends.hpp" -#include "kmm/core/reduction.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -/** - * - */ -void execute_gpu_reduction_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - ReductionDef reduction -); - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/memops/host_copy.hpp b/include/kmm/memops/host_copy.hpp deleted file mode 100644 index 96c4c4b4..00000000 --- a/include/kmm/memops/host_copy.hpp +++ /dev/null @@ -1,9 +0,0 @@ -#pragma once - -#include "kmm/memops/types.hpp" - -namespace kmm { - -void execute_copy(const void* src_buffer, void* dst_buffer, CopyDef copy_def); - -} \ No newline at end of file diff --git a/include/kmm/memops/host_fill.hpp b/include/kmm/memops/host_fill.hpp deleted file mode 100644 index b0992811..00000000 --- a/include/kmm/memops/host_fill.hpp +++ /dev/null @@ -1,9 +0,0 @@ -#include - -#include "kmm/memops/types.hpp" - -namespace kmm { - -void execute_fill(void* dst_buffer, const FillDef& fill); - -} \ No newline at end of file diff --git a/include/kmm/memops/host_reduction.hpp b/include/kmm/memops/host_reduction.hpp deleted file mode 100644 index 6d71bff5..00000000 --- a/include/kmm/memops/host_reduction.hpp +++ /dev/null @@ -1,12 +0,0 @@ -#pragma once - -#include "kmm/memops/types.hpp" - -namespace kmm { - -/** - * - */ -void execute_reduction(const void* src_buffer, void* dst_buffer, ReductionDef reduction); - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/memops/types.hpp b/include/kmm/memops/types.hpp deleted file mode 100644 index 8e4a2f7f..00000000 --- a/include/kmm/memops/types.hpp +++ /dev/null @@ -1,73 +0,0 @@ -#pragma once - -#include "kmm/core/reduction.hpp" -#include "kmm/utils/small_vector.hpp" - -namespace kmm { - -struct CopyDef { - static constexpr size_t MAX_DIMS = 3; - - CopyDef(size_t element_size = 0) : element_size(element_size) {} - - void add_dimension(size_t count, size_t src_offset, size_t dst_offset); - - void add_dimension( - size_t count, - size_t src_offset, - size_t dst_offset, - size_t src_stride, - size_t dst_stride - ); - - size_t minimum_source_bytes_needed() const; - size_t minimum_destination_bytes_needed() const; - size_t number_of_bytes_copied() const; - size_t effective_dimensionality() const; - - void simplify(); - - size_t element_size = 0; - size_t src_offset = 0; - size_t dst_offset = 0; - size_t counts[MAX_DIMS] = {1, 1, 1}; - size_t src_strides[MAX_DIMS] = {0, 0, 0}; - size_t dst_strides[MAX_DIMS] = {0, 0, 0}; -}; - -struct FillDef { - FillDef(size_t element_length, size_t num_elements, const void* fill_value) : - offset_elements(0), - num_elements(num_elements) { - this->fill_value.insert_all( - reinterpret_cast(fill_value), - reinterpret_cast(fill_value) + element_length - ); - } - - template - static FillDef with_value(const T& value, size_t num_elements = 1) { - return {sizeof(T), num_elements, &value}; - } - - size_t minimum_destination_bytes_needed() const; - - size_t offset_elements = 0; - size_t num_elements; - byte_buffer fill_value; -}; - -struct ReductionDef { - size_t minimum_source_bytes_needed() const; - size_t minimum_destination_bytes_needed() const; - - Reduction operation; - DataType data_type; - size_t num_outputs; - size_t num_inputs_per_output = 1; - size_t input_stride_elements = num_outputs; - size_t input_offset_elements = 0; - size_t output_offset_elements = 0; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/planner/array_descriptor.hpp b/include/kmm/planner/array_descriptor.hpp deleted file mode 100644 index 2b79670f..00000000 --- a/include/kmm/planner/array_descriptor.hpp +++ /dev/null @@ -1,50 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "kmm/core/distribution.hpp" -#include "kmm/core/identifiers.hpp" -#include "kmm/utils/geometry.hpp" - -namespace kmm { - -class TaskGraph; - -struct BufferDescriptor { - BufferId id; - BufferLayout layout; - EventId last_write_event {}; - EventList last_access_events {}; -}; - -template -class ArrayDescriptor { - KMM_NOT_COPYABLE_OR_MOVABLE(ArrayDescriptor) - - public: - ArrayDescriptor(TaskGraph& stage, Distribution distribution, DataType dtype); - - const Distribution& distribution() const { - return m_distribution; - } - - DataType data_type() const { - return m_dtype; - } - - EventId copy_bytes_into_buffer(TaskGraph& stage, void* dst_data); - EventId copy_bytes_from_buffer(TaskGraph& stage, const void* src_data); - - EventId join_events(TaskGraph& stage) const; - void destroy(TaskGraph& stage); - - public: - mutable std::shared_mutex m_mutex; - Distribution m_distribution; - DataType m_dtype; - std::vector m_buffers; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/planner/read_planner.hpp b/include/kmm/planner/read_planner.hpp deleted file mode 100644 index fb25fe9b..00000000 --- a/include/kmm/planner/read_planner.hpp +++ /dev/null @@ -1,32 +0,0 @@ -#pragma once - -#include "kmm/planner/array_descriptor.hpp" - -namespace kmm { - -template -class ArrayReadPlanner { - KMM_NOT_COPYABLE_OR_MOVABLE(ArrayReadPlanner) - - public: - ArrayReadPlanner(std::shared_ptr> instance); - ~ArrayReadPlanner(); - - BufferRequirement prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - EventList& deps_out - ); - - void finalize_access(TaskGraph& stage, EventId event_id); - - void commit(TaskGraph& stage); - - private: - std::shared_lock m_lock; - std::shared_ptr> m_instance; - std::vector> m_read_events; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/planner/reduction_planner.hpp b/include/kmm/planner/reduction_planner.hpp deleted file mode 100644 index db450e50..00000000 --- a/include/kmm/planner/reduction_planner.hpp +++ /dev/null @@ -1,59 +0,0 @@ -#pragma once - -#include "kmm/core/reduction.hpp" -#include "kmm/planner/array_descriptor.hpp" - -namespace kmm { - -struct PartialReductionBuffer { - size_t chunk_index; - BufferId buffer_id; - MemoryId memory_id; - size_t replication_factor; - EventId creation_event; - EventList write_events; -}; - -template -class ArrayReductionPlanner { - KMM_NOT_COPYABLE_OR_MOVABLE(ArrayReductionPlanner) - - public: - ArrayReductionPlanner(std::shared_ptr> instance, Reduction op); - ~ArrayReductionPlanner(); - - BufferRequirement prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - size_t replication_factor, - EventList& deps_out - ); - - void finalize_access(TaskGraph& stage, EventId event_id); - - void commit(TaskGraph& stage); - - private: - EventId reduce_per_chunk( - TaskGraph& stage, - size_t chunk_index, - PartialReductionBuffer** buffers, - size_t num_buffers - ); - - std::pair reduce_per_chunk_and_memory( - TaskGraph& stage, - size_t chunk_index, - MemoryId memory_id, - PartialReductionBuffer** buffers, - size_t num_buffers - ); - - std::unique_lock m_lock; - std::shared_ptr> m_instance; - std::vector m_partial_buffers; - Reduction m_reduction; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/planner/write_planner.hpp b/include/kmm/planner/write_planner.hpp deleted file mode 100644 index e4f217fe..00000000 --- a/include/kmm/planner/write_planner.hpp +++ /dev/null @@ -1,31 +0,0 @@ -#pragma once - -#include "kmm/planner/array_descriptor.hpp" - -namespace kmm { - -template -class ArrayWritePlanner { - KMM_NOT_COPYABLE_OR_MOVABLE(ArrayWritePlanner) - - public: - ArrayWritePlanner(std::shared_ptr> instance); - ~ArrayWritePlanner(); - - BufferRequirement prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - EventList& deps_out - ); - - void finalize_access(TaskGraph& stage, EventId event_id); - - void commit(TaskGraph& stage); - - private: - std::unique_lock m_lock; - std::shared_ptr> m_instance; - std::vector> m_write_events; -}; -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/allocators/arena.hpp b/include/kmm/runtime/allocators/arena.hpp new file mode 100644 index 00000000..49b903c0 --- /dev/null +++ b/include/kmm/runtime/allocators/arena.hpp @@ -0,0 +1,102 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +#include "kmm/runtime/allocators/base.hpp" + +namespace kmm { + +class ArenaAllocator: public Allocator { + KMM_NOT_COPYABLE_OR_MOVABLE(ArenaAllocator) + + struct Chunk { + size_t size; + DeviceEventSet deps; + }; + + struct Block { + void* base; + size_t size; + + // Free chunks of this block. Indexed by offset (for coalescing with other free chunks) + // and by (size, offset) (for finding available space). + std::map by_offset; + std::set> by_size; + }; + + struct Allocation { + Block* block; + size_t offset; + size_t size; + }; + + public: + ArenaAllocator(std::unique_ptr inner, size_t block_size = 512UL << 20); + ~ArenaAllocator(); + + AllocResult allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out + ) override final; + + void deallocate_async( + const DeviceStream& stream, + void* addr, + BufferLayout layout + ) override final; + + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + + void deallocate(void* addr, BufferLayout layout) override final; + + void poll() override final; + void trim(size_t nbytes_remaining) override final; + bool trim_one(const DeviceStream* stream_opt); + + // Total number of bytes reserved from the inner allocator (i.e. the sum of block sizes). + std::optional bytes_reserved() const override final { + return m_bytes_reserved; + } + + private: + AllocResult allocate_generic( + const DeviceStream* stream_opt, + BufferLayout layout, + void** addr_out + ); + + void deallocate_generic(const DeviceStream* stream_opt, void* addr, BufferLayout layout); + + AllocResult add_block(const DeviceStream* stream, size_t min_size); + + bool find_best_fit( + const DeviceStream* stream, + size_t nbytes, + Block*& block_out, + size_t& offset_out + ) const; + + static void insert_free(Block& block, size_t offset, size_t size, DeviceEventSet dep); + static Chunk take_free(Block& block, size_t offset); + + std::unique_ptr m_inner; + DeviceEventRegistry m_events; + size_t m_block_size; + size_t m_bytes_reserved = 0; + std::vector> m_blocks; + std::unordered_map m_allocations; + + // Index of the block where the previous search succeeded. Just a search-order + // hint (taken modulo `m_blocks.size()`), so it stays harmless across trims. + mutable size_t m_search_hint = 0; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/allocators/base.hpp b/include/kmm/runtime/allocators/base.hpp index b22eacad..370a2761 100644 --- a/include/kmm/runtime/allocators/base.hpp +++ b/include/kmm/runtime/allocators/base.hpp @@ -1,98 +1,60 @@ #pragma once #include -#include +#include +#include -#include "kmm/runtime/stream_manager.hpp" -#include "kmm/utils/macros.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/device_stream.hpp" namespace kmm { +enum struct AllocResult { Success, ErrorOutOfMemory, ErrorUnsupported, ErrorPending }; -enum struct AllocationResult { - Success, - ErrorOutOfMemory, -}; +class Allocator { + KMM_NOT_COPYABLE_OR_MOVABLE(Allocator) -/** - * Abstract base class for all asynchronous memory allocators. - */ -class AsyncAllocator { public: - virtual ~AsyncAllocator() = default; - - /** - * Allocates `nbytes` of memory and returns the pointer in `addr_out`. The allocated - * region can only be used after the events in `deps_out` have completed. - * - * Returns `AllocationResult::Success` if the operation was successful. - */ - virtual AllocationResult allocate_async( - size_t nbytes, - void** addr_out, - DeviceEventSet& deps_out - ) = 0; - - /** - * Deallocates the give address. This address MUST be previously allocated using - * `allocated_async` with the exact same size. The `deps` parameter can be used to specify any - * dependencies that must be satisfied e the memory is actually deallocated - */ - virtual void deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps = {}) = 0; - - /** - * Perform any pending asynchronous operations. - */ - virtual void make_progress() {} - - /** - * Trim unused memory to reduce the allocator's footprint. - */ - virtual void trim(size_t nbytes_remaining = 0) {} -}; -/** - * Abstract base class for all synchronous memory allocators. - */ -class SyncAllocator: public AsyncAllocator { - KMM_NOT_COPYABLE_OR_MOVABLE(SyncAllocator) - - public: - SyncAllocator( - std::shared_ptr streams, - size_t max_bytes = std::numeric_limits::max() + Allocator(); + virtual ~Allocator(); + + /// allocate memory asynchronously on the given stream. + virtual AllocResult allocate_async( // + const DeviceStream& stream, + BufferLayout layout, + void** addr_out ); - ~SyncAllocator(); + /// deallocate memory asynchronously on the given stream. The memory must previously be allocated + /// using `allocate_async` or `allocate`. + virtual void deallocate_async( // + const DeviceStream& stream, + void* addr, + BufferLayout layout + ); - /** - * Allocate `nbytes` bytes of memory and sets the pointer in `addr_out`. - * - * Returns `AllocationResult::Success` if the operation was successful. - */ - virtual AllocationResult allocate(size_t nbytes, void** addr_out) = 0; + /// allocate memory synchronously, potentially blocking until memory becomes available. + virtual AllocResult allocate(BufferLayout layout, void** addr_out) = 0; - /** - * Deallocates the give address. This address MUST be previously allocated using - * `allocate` with the exact same size. - */ - virtual void deallocate(void* addr, size_t nbytes) = 0; + /// deallocate memory synchronously. The memory must previously be allocated using `allocate_async` or `allocate`. + virtual void deallocate(void* addr, BufferLayout layout) = 0; - AllocationResult allocate_async(size_t nbytes, void** addr_out, DeviceEventSet& deps_out) final; - void deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) final; - void make_progress() final; - void trim(size_t nbytes_remaining = 0) final; + /// Called many times per second, allowing the allocator to update internal bookkeeping. + virtual void poll() {} - private: - struct DeferredDealloc { - void* addr; - size_t nbytes; - DeviceEventSet dependencies; - }; + /// Reduce the number of bytes this allocator holds reserved from the OS/driver to the given limit. For example, + /// for a block/pool allocator, this will free unused blocks until only the given number of bytes remain. + /// Note that this method is a hint as allocators may sometimes not be able to trim to the given limit. + virtual void trim(size_t nbytes_remaining) {} - std::shared_ptr m_streams; - std::deque m_pending_deallocs; - size_t m_bytes_limit; - size_t m_bytes_in_use = 0; + /// Real number of bytes this allocator holds reserved from the OS/driver (which, for a + /// block/pool allocator, is much larger than the sum of the live allocation sizes). + /// `std::nullopt` if the allocator does not track this. + virtual std::optional bytes_reserved() const { + return std::nullopt; + } }; -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/runtime/allocators/block.hpp b/include/kmm/runtime/allocators/block.hpp deleted file mode 100644 index 48a48425..00000000 --- a/include/kmm/runtime/allocators/block.hpp +++ /dev/null @@ -1,56 +0,0 @@ -#pragma once - -#include -#include -#include -#include - -#include "kmm/runtime/allocators/base.hpp" - -namespace kmm { - -class BlockAllocator: public AsyncAllocator { - public: - static constexpr size_t DEFAULT_BLOCK_SIZE = 1024L * 1024 * 500; - - BlockAllocator( - std::unique_ptr allocator, - size_t min_block_size = DEFAULT_BLOCK_SIZE - ); - ~BlockAllocator(); - AllocationResult allocate_async(size_t nbytes, void** addr_out, DeviceEventSet& deps_out) final; - void deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) final; - void make_progress() final; - void trim(size_t nbytes_remaining = 0) final; - - private: - struct Block; - struct BlockRegion; - struct BlockRegionSize { - size_t value; - }; - - struct RegionSizeCompare { - using is_transparent = void; - bool operator()(const BlockRegion*, const BlockRegion*) const; - bool operator()(const BlockRegion*, BlockRegionSize) const; - }; - - BlockRegion* allocate_block(size_t min_nbytes); - BlockRegion* find_region(size_t nbytes, size_t alignment); - static std::pair split_region( - BlockRegion* region, - size_t left_size - ); - static BlockRegion* merge_regions(BlockRegion* left, BlockRegion* right); - static size_t offset_to_alignment(const BlockRegion* region, size_t alignment); - static bool fits_in_region(const BlockRegion* region, size_t nbytes, size_t alignment); - - std::unique_ptr m_allocator; - std::unordered_map m_active_regions; - std::vector> m_blocks; - size_t m_active_block = 0; - size_t m_min_block_size = DEFAULT_BLOCK_SIZE; - size_t m_bytes_allocated = 0; -}; -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/allocators/caching.hpp b/include/kmm/runtime/allocators/caching.hpp deleted file mode 100644 index b0ab9ebf..00000000 --- a/include/kmm/runtime/allocators/caching.hpp +++ /dev/null @@ -1,39 +0,0 @@ -#pragma once - -#include -#include - -#include "kmm/runtime/allocators/base.hpp" - -namespace kmm { - -class CachingAllocator: public AsyncAllocator { - public: - CachingAllocator(std::unique_ptr allocator, double max_fragmentation=0.75, size_t initial_watermark = 0); - ~CachingAllocator(); - AllocationResult allocate_async(size_t nbytes, void** addr_out, DeviceEventSet& deps_out) final; - void deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) final; - void make_progress() final; - void trim(size_t nbytes_remaining = 0) final; - size_t free_some_memory(); - - private: - bool can_allocate_bytes(size_t nbytes) const; - - struct AllocationSlot; - struct AllocationBin { - std::unique_ptr head; - AllocationSlot* tail; - }; - - std::unique_ptr m_allocator; - std::unordered_map m_allocation_bins; - AllocationSlot* m_lru_oldest = nullptr; - AllocationSlot* m_lru_newest = nullptr; - size_t m_bytes_watermark = 0; - size_t m_bytes_in_use = 0; - size_t m_bytes_allocated = 0; - double m_max_fragmentation = 0.75; -}; - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/allocators/device.hpp b/include/kmm/runtime/allocators/device.hpp index 682d8bdd..c8485d78 100644 --- a/include/kmm/runtime/allocators/device.hpp +++ b/include/kmm/runtime/allocators/device.hpp @@ -4,70 +4,15 @@ namespace kmm { -struct Allocation { - void* addr; - size_t nbytes; - DeviceEvent event; -}; - -class PinnedMemoryAllocator: public SyncAllocator { - public: - PinnedMemoryAllocator( - GPUContextHandle context, - std::shared_ptr streams, - size_t max_bytes = std::numeric_limits::max() - ); - - AllocationResult allocate(size_t nbytes, void** addr_out) final; - void deallocate(void* addr, size_t nbytes) final; - - private: - GPUContextHandle m_context; -}; - -class DeviceMemoryAllocator: public SyncAllocator { +class DeviceMemoryAllocator: public Allocator { public: - DeviceMemoryAllocator( - GPUContextHandle context, - std::shared_ptr streams, - size_t max_bytes = std::numeric_limits::max() - ); + DeviceMemoryAllocator(g_context_t context); - AllocationResult allocate(size_t nbytes, void** addr_out) final; - void deallocate(void* addr, size_t nbytes) final; - - private: - GPUContextHandle m_context; -}; - -enum struct DevicePoolKind { Default, Create }; - -class DevicePoolAllocator: public AsyncAllocator { - KMM_NOT_COPYABLE_OR_MOVABLE(DevicePoolAllocator) - - public: - DevicePoolAllocator( - GPUContextHandle context, - std::shared_ptr streams, - DevicePoolKind kind = DevicePoolKind::Create, - size_t max_bytes = std::numeric_limits::max() - ); - ~DevicePoolAllocator(); - AllocationResult allocate_async(size_t nbytes, void** addr_out, DeviceEventSet& deps_out) final; - void deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) final; - void make_progress() final; - void trim(size_t nbytes_remaining = 0) final; + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + void deallocate(void* addr, BufferLayout layout) override final; private: - GPUContextHandle m_context; - g_memory_pool_t m_pool; - std::shared_ptr m_streams; - DeviceStream m_alloc_stream; - DeviceStream m_dealloc_stream; - std::deque m_pending_deallocs; - DevicePoolKind m_kind; - size_t m_bytes_in_use = 0; - size_t m_bytes_limit; + g_context_t m_context; }; -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/runtime/allocators/device_pool.hpp b/include/kmm/runtime/allocators/device_pool.hpp new file mode 100644 index 00000000..60a48af8 --- /dev/null +++ b/include/kmm/runtime/allocators/device_pool.hpp @@ -0,0 +1,48 @@ +#pragma once + +#include + +#include "kmm/runtime/allocators/base.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +enum struct DevicePoolKind { Default, Create }; + +class DevicePoolAllocator: public Allocator { + KMM_NOT_COPYABLE_OR_MOVABLE(DevicePoolAllocator) + + public: + DevicePoolAllocator( + g_context_t context, + DevicePoolKind kind = DevicePoolKind::Create, + size_t max_size = std::numeric_limits::max() + ); + ~DevicePoolAllocator(); + + AllocResult allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out + ) override final; + + void deallocate_async( + const DeviceStream& stream, + void* addr, + BufferLayout layout + ) override final; + + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + + void deallocate(void* addr, BufferLayout layout) override final; + + void trim(size_t nbytes_remaining) override final; + + private: + g_context_t m_context; + g_memory_pool_t m_pool; + DevicePoolKind m_kind; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/allocators/limit.hpp b/include/kmm/runtime/allocators/limit.hpp new file mode 100644 index 00000000..732b3375 --- /dev/null +++ b/include/kmm/runtime/allocators/limit.hpp @@ -0,0 +1,61 @@ +#pragma once + +#include "kmm/runtime/allocators/base.hpp" + +namespace kmm { + +class LimitAllocator: public Allocator { + KMM_NOT_COPYABLE_OR_MOVABLE(LimitAllocator) + + struct Allocation { + void* addr; + size_t nbytes; + DeviceEvent event; + }; + + public: + LimitAllocator(std::unique_ptr inner, DeviceEventRegistry events, size_t max_size); + ~LimitAllocator(); + + AllocResult allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out + ) override final; + + void deallocate_async( // + const DeviceStream& stream, + void* addr, + BufferLayout layout + ) override final; + + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + + void deallocate(void* addr, BufferLayout layout) override final; + + void poll() final; + + void trim(size_t nbytes_remaining) final; + + std::optional bytes_reserved() const override final; + + private: + bool ensure_enough_space(const DeviceStream* stream, size_t nbytes); + + std::unique_ptr m_inner; + DeviceEventRegistry m_events; + std::deque m_pending_deallocs; + DeviceEventSet m_limit_barrier; + + // maximum number of bytes that can be allocated + size_t m_bytes_limit = 0; + + // number of bytes in use. For these, allocate has been called but not deallocate yet. + size_t m_bytes_active = 0; + + // number of bytes awaiting deallocation. For these, deallocate hase been called and + // the deallocation event is current waiting in m_pending_deallocs. + size_t m_bytes_pending = 0; +}; + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/allocators/managed.hpp b/include/kmm/runtime/allocators/managed.hpp new file mode 100644 index 00000000..f1ed7cd9 --- /dev/null +++ b/include/kmm/runtime/allocators/managed.hpp @@ -0,0 +1,18 @@ +#pragma once + +#include "kmm/runtime/allocators/base.hpp" + +namespace kmm { + +class ManagedMemoryAllocator: public Allocator { + public: + ManagedMemoryAllocator(g_context_t context); + + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + void deallocate(void* addr, BufferLayout layout) override final; + + private: + g_context_t m_context; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/allocators/pinned.hpp b/include/kmm/runtime/allocators/pinned.hpp new file mode 100644 index 00000000..f304329a --- /dev/null +++ b/include/kmm/runtime/allocators/pinned.hpp @@ -0,0 +1,18 @@ +#pragma once + +#include "kmm/runtime/allocators/base.hpp" + +namespace kmm { + +class PinnedMemoryAllocator: public Allocator { + public: + PinnedMemoryAllocator(g_context_t context); + + AllocResult allocate(BufferLayout layout, void** addr_out) override final; + void deallocate(void* addr, BufferLayout layout) override final; + + private: + g_context_t m_context; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/allocators/system.hpp b/include/kmm/runtime/allocators/system.hpp index 720c91d8..8857dd41 100644 --- a/include/kmm/runtime/allocators/system.hpp +++ b/include/kmm/runtime/allocators/system.hpp @@ -4,17 +4,9 @@ namespace kmm { -class SystemAllocator: public SyncAllocator { - public: - SystemAllocator( - std::shared_ptr streams, - size_t max_bytes = std::numeric_limits::max() - ) : - SyncAllocator(streams, max_bytes) {} - - protected: - AllocationResult allocate(size_t nbytes, void** addr_out) final; - void deallocate(void* addr, size_t nbytes) final; +class SystemAllocator: public Allocator { + AllocResult allocate(BufferLayout layout, void** addr_out); + void deallocate(void* addr, BufferLayout layout); }; } // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/buffer.hpp b/include/kmm/runtime/buffer.hpp new file mode 100644 index 00000000..2a7a2f1c --- /dev/null +++ b/include/kmm/runtime/buffer.hpp @@ -0,0 +1,45 @@ +#pragma once + +#include +#include + +#include "kmm/runtime/identifiers.hpp" +#include "kmm/utils/refcnt_ptr.hpp" + +namespace kmm { + +enum struct AccessMode { + /// multiple requests may be active concurrently in multiple memories, but they cannot write + Read, + + /// only one request may be active in one memory + ReadWrite, + + /// multiple access will be reduced into one + Reduce +}; + +struct BufferLayout { + BufferLayout repeat(size_t n) { + size_t remainder = size_in_bytes % alignment; + size_t padding = remainder != 0 ? alignment - remainder : 0; + return {(size_in_bytes + padding) * n, alignment}; + } + + template + static BufferLayout for_type(size_t n = 1) { + return BufferLayout {sizeof(T), alignof(T)}.repeat(n); + } + + size_t size_in_bytes = 0; + size_t alignment = 1; +}; + +struct BufferAccessor { + MemoryId memory_id = MemoryId::host(); + size_t size_in_bytes = 0; + bool is_writable = false; + void* address = nullptr; +}; + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/buffer_registry.hpp b/include/kmm/runtime/buffer_registry.hpp deleted file mode 100644 index d9bb6f32..00000000 --- a/include/kmm/runtime/buffer_registry.hpp +++ /dev/null @@ -1,57 +0,0 @@ -#pragma once - -#include -#include -#include -#include - -#include "kmm/runtime/memory_manager.hpp" - -namespace kmm { - -class PoisonException; -using BufferRequest = std::shared_ptr; -using BufferRequestList = std::vector; - -class BufferRegistry { - public: - BufferRegistry(std::shared_ptr memory_manager); - - BufferId add(BufferId id, BufferLayout layout); - - std::shared_ptr get(BufferId id); - - void remove(BufferId buffer_id); - - BufferRequestList create_requests(const std::vector& buffers); - - Poll poll_requests(const BufferRequestList& requests, DeviceEventSet& dependencies_out); - - std::vector access_requests(const BufferRequestList& requests); - - void release_requests(BufferRequestList& requests, DeviceEvent event = {}); - - void poison(BufferId id, PoisonException reason); - - void poison_all(const std::vector& buffers, PoisonException reason); - - private: - struct BufferMeta { - std::shared_ptr buffer; - std::unique_ptr poison_reason_opt = nullptr; - }; - - std::shared_ptr m_memory_manager; - std::unordered_map m_buffers; -}; - -class PoisonException final: public std::exception { - public: - PoisonException(const std::string& error); - const char* what() const noexcept override; - - private: - std::string m_message; -}; - -} // namespace kmm diff --git a/include/kmm/runtime/data_interfaces/base.hpp b/include/kmm/runtime/data_interfaces/base.hpp new file mode 100644 index 00000000..f5255fed --- /dev/null +++ b/include/kmm/runtime/data_interfaces/base.hpp @@ -0,0 +1,121 @@ +#pragma once + +#include + +#include "kmm/runtime/allocators/base.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/identifiers.hpp" + +namespace kmm { + +class MemorySystem; + +/// Per-buffer counterpart to `MemorySystem`: knows the shape of a single buffer and how to +/// allocate, deallocate, and copy exactly that buffer at any of its possible locations. The +/// `MemorySystem` that backs the buffer is not held by the interface; the caller passes it into +/// every method that needs it. +/// +/// Concrete implementations live in their own headers under `data_interfaces/`: `FlatDataInterface` +/// (flat.hpp), `ManagedDataInterface` (managed.hpp), `PinnedDataInterface` (pinned.hpp), and +/// `ExternalDataInterface` (external.hpp). +class DataInterface { + public: + virtual ~DataInterface() = default; + + virtual size_t size_in_bytes() const noexcept = 0; + + /// Allocate this buffer on the given memory id. The returned dependencies indicate + /// when the buffer is safe to be used. + virtual AllocResult allocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ) = 0; + + /// Deallocate this buffer on the given memory id. The given dependencies indicate when + /// the deallocation should happen. + virtual void deallocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) = 0; + + /// Returns the pointer to the data for the given MemoryId. The caller must ensure that + /// `allocate` was called before this function was called. + virtual void* address(MemoryId memory_id) const noexcept = 0; + + /// Copy data from the given `src` to `dst` memory. The given dependencies indicate + /// when the copy should be performed. The caller must have called `allocate` on both + /// the `src` and `dst` in the past. + virtual void copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) = 0; + + /// Indicate if copying data from `src` to `dst` is supported. + virtual bool is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst + ) const noexcept = 0; + + /// Hint that this buffer will be accessed on the given memory in the future when the given + /// dependencies complete. This is just a hint, useful for `cudaMemPrefetchAsync`. + virtual void hint_access( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) {} + + /// Initialize the buffer on the host. This returns a future as it is likely that this + /// operation will be performed on the CPU. + virtual std::future initialize_host(MemorySystem& system, const DeviceEventSet& deps) { + return {}; + } + + /// Initialize the buffer on the GPU. + virtual DeviceEvent initialize_device( + MemorySystem& system, + DeviceId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) { + return DeviceEvent::null(); + } + + /// Allocate the buffer in `dst` memory and copy data immediately from `src` to `dst. + /// This is typically just an alias for calling `allocate` followed by `copy`, but might + /// be optimized + virtual AllocResult allocate_and_copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) { + DeviceEventSet deps; + auto result = allocate(system, dst, stream_hint, deps); + + if (result == AllocResult::Success) { + try { + deps.insert(deps_in); + copy(system, src, dst, stream_hint, deps, deps_out); + } catch (...) { + deallocate(system, dst, stream_hint, deps); + throw; + } + } + + return result; + } +}; + +} // namespace kmm diff --git a/include/kmm/runtime/data_interfaces/external.hpp b/include/kmm/runtime/data_interfaces/external.hpp new file mode 100644 index 00000000..04d5d615 --- /dev/null +++ b/include/kmm/runtime/data_interfaces/external.hpp @@ -0,0 +1,58 @@ +#pragma once + +#include "kmm/runtime/data_interfaces/base.hpp" + +namespace kmm { + +/// A `DataInterface` that wraps a pre-existing, externally-owned pointer pinned to a single +/// `MemoryId`. KMM never allocates, deallocates, or copies the underlying memory: `allocate` +/// and `deallocate` on the pinned memory are no-ops, and any attempt to access or copy the +/// buffer on a different `MemoryId` throws `std::runtime_error`. +class ExternalDataInterface final: public DataInterface { + public: + ExternalDataInterface(void* ptr, size_t size_in_bytes, MemoryId memory_id); + + size_t size_in_bytes() const noexcept override; + + AllocResult allocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ) override; + + void deallocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + void* address( // + MemoryId memory_id + ) const noexcept override; + + void copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) override; + + bool is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst + ) const noexcept override; + + private: + void check_memory_id(MemoryId memory_id) const; + + void* m_ptr; + size_t m_size_in_bytes; + MemoryId m_memory_id; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/data_interfaces/flat.hpp b/include/kmm/runtime/data_interfaces/flat.hpp new file mode 100644 index 00000000..e295ad59 --- /dev/null +++ b/include/kmm/runtime/data_interfaces/flat.hpp @@ -0,0 +1,79 @@ +#pragma once + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/data_interfaces/base.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +/// Default `DataInterface`: a flat buffer of `layout.size_in_bytes` bytes, allocated and copied +/// through a shared `MemorySystem`. Equivalent to how every buffer behaved before per-buffer +/// `DataInterface`s existed. +class FlatDataInterface final: public DataInterface { + public: + /// If `fill_value` is non-empty, the buffer is filled with copies of it the first time it is + /// materialized in any memory (see `initialize_host`/`initialize_device`). + FlatDataInterface(BufferLayout layout, FillValue fill_value = {}); + + size_t size_in_bytes() const noexcept override; + + AllocResult allocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ) override; + + void deallocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + void* address( // + MemoryId memory_id + ) const noexcept override; + + bool is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst + ) const noexcept override; + + void copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) override; + + std::future initialize_host(MemorySystem& system, const DeviceEventSet& deps) override; + + DeviceEvent initialize_device( + MemorySystem& system, + DeviceId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + AllocResult allocate_and_copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) override; + + private: + BufferLayout m_layout; + FillValue m_fill_value; + void* m_host_ptr = nullptr; + g_device_ptr_t m_device_ptrs[MAX_DEVICES] {}; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/data_interfaces/managed.hpp b/include/kmm/runtime/data_interfaces/managed.hpp new file mode 100644 index 00000000..96f43191 --- /dev/null +++ b/include/kmm/runtime/data_interfaces/managed.hpp @@ -0,0 +1,85 @@ +#pragma once + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/data_interfaces/base.hpp" +#include "kmm/runtime/memops/fill.hpp" + +namespace kmm { + +/// A `DataInterface` backed by a single CUDA/HIP managed-memory allocation (`cudaMallocManaged`), +/// resident on the host and every device simultaneously. There is only ever one physical +/// allocation: `allocate`/`deallocate` are refcounted no-ops beyond the first/last call +/// (regardless of which `MemoryId` is requested), `address` returns the same pointer for every +/// `MemoryId`, and `copy` is a no-op since the driver keeps the allocation coherent everywhere. +class ManagedDataInterface final: public DataInterface { + public: + /// If `fill_value` is non-empty, the buffer is filled with copies of it the first time it is + /// materialized (see `initialize_host`/`initialize_device`). + ManagedDataInterface(BufferLayout layout, FillValue fill_value = {}); + + size_t size_in_bytes() const noexcept override; + + AllocResult allocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ) override; + + void deallocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + void* address( // + MemoryId memory_id + ) const noexcept override; + + bool is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst + ) const noexcept override; + + /// Prefetches the allocation to `memory_id` (`cudaMemPrefetchAsync`/`hipMemPrefetchAsync`) so + /// the driver can start migrating pages before the actual access. Best-effort: fired on the + /// default stream without waiting on `deps`, since a prefetch racing with a pending write is + /// merely a missed optimization (the driver still page-faults correctly), not a correctness + /// issue. + void hint_access( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + void copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) override; + + std::future initialize_host(MemorySystem& system, const DeviceEventSet& deps) override; + + DeviceEvent initialize_device( + MemorySystem& system, + DeviceId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + private: + BufferLayout m_layout; + FillValue m_fill_value; + void* m_ptr = nullptr; + size_t m_refcount = 0; + DeviceEventSet m_alloc_deps; + DeviceEventSet m_dealloc_deps; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/data_interfaces/pinned.hpp b/include/kmm/runtime/data_interfaces/pinned.hpp new file mode 100644 index 00000000..89444eeb --- /dev/null +++ b/include/kmm/runtime/data_interfaces/pinned.hpp @@ -0,0 +1,61 @@ +#pragma once + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/data_interfaces/base.hpp" +#include "kmm/runtime/device_event_registry.hpp" +#include "kmm/runtime/memops/fill.hpp" + +namespace kmm { + +/// A `DataInterface` backed by a single pinned host allocation. Devices access it directly +/// (zero-copy) through a mapped pointer, so `copy` is a no-op and no device-side copy is ever +/// staged. +class PinnedDataInterface final: public DataInterface { + public: + PinnedDataInterface(BufferLayout layout); + + size_t size_in_bytes() const noexcept override; + + AllocResult allocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ) override; + + void deallocate( // + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ) override; + + void* address( // + MemoryId memory_id + ) const noexcept override; + + bool is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst + ) const noexcept override; + + void copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out + ) override; + + private: + BufferLayout m_layout; + void* m_host_ptr = nullptr; + void* m_device_ptrs[MAX_DEVICES] {}; + size_t m_refcount = 0; + DeviceEventSet m_alloc_deps; + DeviceEventSet m_dealloc_deps; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/device_data_streams.hpp b/include/kmm/runtime/device_data_streams.hpp new file mode 100644 index 00000000..0c9f57d3 --- /dev/null +++ b/include/kmm/runtime/device_data_streams.hpp @@ -0,0 +1,73 @@ +#pragma once + +#include +#include +#include + +#include "kmm/core/macros.hpp" +#include "kmm/runtime/device_event_registry.hpp" +#include "kmm/runtime/device_stream.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/system_info.hpp" + +namespace kmm { + +enum class StreamKind : uint8_t { DeviceToDevice = 0, HostToDevice = 1, DeviceToHost = 2 }; + +std::ostream& operator<<(std::ostream&, StreamKind kind); + +class DeviceDataStreams { + public: + DeviceDataStreams( + const SystemInfo& info, + DeviceEventRegistry events, + size_t num_d2d_streams, + size_t num_h2d_streams, + size_t num_d2h_streams + ); + + DeviceDataStreams(DeviceDataStreams&&) noexcept; + ~DeviceDataStreams(); + + /** + * Acquire a free stream of the given kind on `device_id`, preferring the one predicted to + * become ready soonest for work depending on `deps`. + */ + DeviceStreamId acquire_stream(DeviceId device_id, StreamKind kind, const DeviceEventSet& deps); + + /** + * See `release_stream(const DeviceStream&, uint64_t)`. + */ + DeviceEvent release_stream(DeviceStreamId stream_id, uint64_t cost = 1); + + /** + * Convenience method that acquires a stream of the given kind on `device_id`, calls + * `fun(stream)` to submit work onto it, and then releases the stream again. + */ + template + DeviceEvent submit(DeviceId device_id, StreamKind kind, const DeviceEventSet& deps, F&& fun) { + auto stream_id = acquire_stream(device_id, kind, deps); + uint64_t cost; + + try { + cost = std::forward(fun)(stream_id); + } catch (...) { + release_stream(stream_id); + throw; + } + + return release_stream(stream_id, cost); + } + + void make_progress(); + + struct Impl; + + private: + std::unique_ptr m_impl; +}; + +} // namespace kmm + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/runtime/device_event.hpp b/include/kmm/runtime/device_event.hpp new file mode 100644 index 00000000..6a4b7fbd --- /dev/null +++ b/include/kmm/runtime/device_event.hpp @@ -0,0 +1,179 @@ +#pragma once + +#include "fmt/format.h" + +#include "kmm/utils/hash_utils.hpp" +#include "kmm/utils/small_vector.hpp" + +namespace kmm { + +static constexpr uint64_t MAX_DEVICE_STREAMS = 256; +static constexpr uint64_t MAX_DEVICE_EVENTS = (~uint64_t(0)) / MAX_DEVICE_STREAMS - 1; +class DeviceEventRegistry; +class PrecedenceVector; + +class DeviceStreamId { + static constexpr uint64_t INVALID_INDEX = ~uint64_t(0); + + public: + DeviceStreamId() noexcept = default; + + explicit DeviceStreamId(uint64_t index) : m_index(index) { + KMM_ASSERT(index < MAX_DEVICE_STREAMS); + } + + static DeviceStreamId null() noexcept { + return {}; + } + + uint64_t get() const noexcept { + KMM_UNSAFE_ASSUME(m_index < MAX_DEVICE_STREAMS); + return m_index; + } + + bool is_null() const noexcept { + return m_index == INVALID_INDEX; + } + + friend std::ostream& operator<<(std::ostream& stream, const DeviceStreamId& e); + + friend bool operator==(const DeviceStreamId& a, const DeviceStreamId& b) { + return a.get() == b.get(); + } + + friend bool operator!=(const DeviceStreamId& a, const DeviceStreamId& b) { + return !(a == b); + } + + private: + uint64_t m_index = INVALID_INDEX; +}; + +class DeviceEvent { + public: + DeviceEvent() noexcept = default; + + DeviceEvent(DeviceStreamId stream_id, uint64_t event_id) { + if (!stream_id.is_null()) { + KMM_ASSERT(stream_id.get() < MAX_DEVICE_STREAMS); + m_event_and_stream_index = stream_id.get() + event_id * MAX_DEVICE_STREAMS; + } + } + + static DeviceEvent null() noexcept { + return {}; + } + + DeviceStreamId stream() const noexcept { + return DeviceStreamId(m_event_and_stream_index % MAX_DEVICE_STREAMS); + } + + uint64_t index() const { + return m_event_and_stream_index / MAX_DEVICE_STREAMS; + } + + size_t hash() const { + return m_event_and_stream_index; + } + + bool is_null() const noexcept { + return m_event_and_stream_index == 0; + } + + bool precedes(const DeviceEvent& that) const noexcept { + return stream() == that.stream() + && this->m_event_and_stream_index <= that.m_event_and_stream_index; + } + + friend std::ostream& operator<<(std::ostream& stream, const DeviceEvent& e); + + friend bool operator==(const DeviceEvent& a, const DeviceEvent& b) { + return a.m_event_and_stream_index == b.m_event_and_stream_index; + } + + friend bool operator<(const DeviceEvent& a, const DeviceEvent& b) { + return a.m_event_and_stream_index < b.m_event_and_stream_index; + } + + friend bool operator<=(const DeviceEvent& a, const DeviceEvent& b) { + return a.m_event_and_stream_index <= b.m_event_and_stream_index; + } + + friend bool operator!=(const DeviceEvent& a, const DeviceEvent& b) { + return !(a == b); + } + + friend bool operator>(const DeviceEvent& a, const DeviceEvent& b) { + return b < a; + } + + friend bool operator>=(const DeviceEvent& a, const DeviceEvent& b) { + return b <= a; + } + + private: + uint64_t m_event_and_stream_index = 0; +}; + +class DeviceEventSet { + public: + DeviceEventSet() noexcept = default; + DeviceEventSet(const DeviceEvent& event); + DeviceEventSet(const DeviceEventSet& events) = default; + DeviceEventSet(DeviceEventSet&& events) noexcept = default; + DeviceEventSet(std::initializer_list list); + + DeviceEventSet& operator=(const DeviceEventSet& that) = default; + DeviceEventSet& operator=(DeviceEventSet&& that) noexcept = default; + DeviceEventSet& operator=(std::initializer_list list); + + void insert(DeviceEvent event) noexcept; + void insert(const DeviceEventSet& events) noexcept; + void insert(DeviceEventSet&& events) noexcept; + void prune(const DeviceEventRegistry& registry) noexcept; + void clear() noexcept; + bool is_empty() const noexcept; + bool contains(const DeviceEvent& event) const noexcept; + bool contains(const DeviceEventSet& events) const noexcept; + DeviceEvent find(DeviceStreamId stream_id) const noexcept; + + const DeviceEvent* begin() const noexcept { + return m_events.begin(); + } + const DeviceEvent* end() const noexcept { + return m_events.end(); + } + + friend std::ostream& operator<<(std::ostream& stream, const DeviceEventSet& e); + + // from device_event_registry.cpp + friend class PrecedenceVector; + + private: + small_vector m_events; +}; + +} // namespace kmm + +template<> +struct std::hash { + size_t operator()(const kmm::DeviceStreamId& id) const noexcept { + return id.get(); + } +}; + +template<> +struct std::hash { + size_t operator()(const kmm::DeviceEvent& id) const noexcept { + return id.hash(); + } +}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/runtime/device_event_registry.hpp b/include/kmm/runtime/device_event_registry.hpp new file mode 100644 index 00000000..2b0f3574 --- /dev/null +++ b/include/kmm/runtime/device_event_registry.hpp @@ -0,0 +1,224 @@ +#pragma once + +#include +#include +#include +#include + +#include "kmm/core/macros.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" +#include "kmm/utils/notify.hpp" +#include "kmm/utils/small_vector.hpp" + +namespace kmm { + +class DeviceStream; + +/** + * The event registry is used to keep track of all GPU events that have been recorded. + */ +class DeviceEventRegistry { + public: + DeviceEventRegistry(); + + /** + * Register a new stream with the event registry. The returned identifier can be used + * to register events or block the stream on existing events. The optional `name` is + * included in log messages that refer to this stream. + */ + DeviceStreamId register_stream(GPUStreamRef stream_ref, std::string name = "") const; + + /** + * Unregister a stream that was previously registered. After this call, no more events + * can be recorded on the stream. + */ + void unregister_stream(DeviceStreamId stream_id) const; + + /** + * Returns the id of the given stream, or `std::nullopt` if it was not registered. + */ + std::optional lookup_stream(GPUStreamId target) const; + + /** + * Returns the id of the given stream, or `std::nullopt` if it was not registered. + */ + DeviceStreamId lookup_or_register_stream(GPUStreamRef stream_ref) const; + + /** + * Calls `unregister_stream` on all registered streams. + */ + void shutdown() const; + + /** + * Get the GPUStream associated with the given stream. + */ + g_stream_t get(DeviceStreamId stream_id) const; + + /** + * Get the GPUContext associated with the given stream. + */ + g_context_t context(DeviceStreamId stream_id) const; + + /** + * Get a `DeviceStream` wrapping the given stream. + */ + DeviceStream stream(DeviceStreamId stream_id) const; + + /** + * Check if the given stream id matches the given context id. + */ + bool has_context(DeviceStreamId stream_id, GPUContextId context_id) const; + + /** + * Record a new event on the given stream. + */ + DeviceEvent record(DeviceStreamId stream_id) const; + + /** + * Convenience method for the common pattern of waiting on dependencies, submitting work + * onto a stream, and recording the resulting event. Waits for `deps`, calls + * `fun(get(stream_id))` to submit work onto the stream, then records and returns the + * resulting event. + */ + template + DeviceEvent submit(DeviceStreamId stream_id, const DeviceEventSet& deps, F&& fun) const { + wait_on_event(stream_id, deps); + std::forward(fun)(get(stream_id)); + return record(stream_id); + } + + /** + * Let the given stream wait until the given event completes. + */ + void wait_on_event(DeviceStreamId stream_id, DeviceEvent event) const; + + /** + * Let the given stream wait until all the given events completes. + */ + void wait_on_event(DeviceStreamId stream_id, const DeviceEventSet& events) const; + + /** + * Let the given stream wait until the given event completes. + */ + void wait_on_event(g_stream_t stream, DeviceEvent event) const; + + /** + * Let the given stream wait until all the given events completes. + */ + void wait_on_event(g_stream_t stream, const DeviceEventSet& events) const; + + /** + * Let the given stream wait until all work already enqueued on the CUDA legacy default + * stream (i.e., stream `0`) in the same context completes. Used to synchronize with + * work launched outside of KMM's own tracked streams. + */ + void wait_on_default_stream(DeviceStreamId stream_id) const; + + /** + * Check if the given stream is currently idle. + */ + bool is_ready(DeviceStreamId stream_id) const; + + /** + * Check if the given event has completed. + */ + bool is_ready(DeviceEvent event) const; + + /** + * Check if all the given events have completed. + */ + bool is_ready(const DeviceEventSet& events) const; + + /** + * Check if all events recorded on all streams have completed. + */ + bool is_all_ready() const; + + /** + * Check if the given event is latest known recorded event on its stream. + */ + bool is_latest(DeviceEvent event) const; + + /** + * Check if one of the given events is latest known recorded event on the given stream id. + */ + bool is_latest_in(DeviceStreamId stream_id, const DeviceEventSet& deps) const; + + /** + * Returns the latest recorded event on the given stream. + */ + DeviceEvent latest_event(DeviceStreamId stream_id) const; + + /** + * Block until the given stream becomes idle. + */ + void synchronize(DeviceStreamId stream_id) const; + + /** + * Block until the given event completes. + */ + void synchronize(DeviceEvent event) const; + + /** + * Block until the given events complete. + */ + void synchronize(const DeviceEventSet& events) const; + + /** + * Block until all events on all streams complete. + */ + void synchronize_all() const; + + /** + * Return a list of events that MUST have completed once the given stream completes. + */ + DeviceEventSet snapshot(DeviceStreamId stream_id) const; + + /** + * Attach a callback that fires when the given event completes. + */ + void attach_callback(DeviceEvent event, NotifyHandle callback) const; + + /** + * Poll the registered streams. + */ + void make_progress() const; + + /** + * Each stream keeps a pool of available events. Once an event completes, the free event + * is added to this pool to be reused again later. + * + * This method clears all these pools and frees these events again. + */ + void trim_event_pool(DeviceStreamId stream_id) const; + + /** + * Calls `trim_event_pool` for all streams. + */ + void trim_event_pool() const; + + /** + * Check if the given event `a` precedes the given stream `b`. In other words, the + * given event MUST complete before the stream becomes idle. This means that there exist + * a dependency between the given event a and the stream b. + * + * Note: this is a hint. If `true`, then `a` MUST precede `b`. However, if `false, then + * it might still be the case but `DeviceEventRegistry` has not witnessed the synchronization. + */ + bool precedes(const DeviceEvent& a, const DeviceStreamId& b) const; + + /** + * Shorthand that checks if `precedes(x, b)` for all x in a. + */ + bool precedes(const DeviceEventSet& a, const DeviceStreamId& b) const; + + struct Impl; + + private: + refcnt_ptr m_impl; +}; + +KMM_REFCNT_TRAITS_FWD(DeviceEventRegistry::Impl) + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/device_resources.hpp b/include/kmm/runtime/device_resources.hpp deleted file mode 100644 index 0550600e..00000000 --- a/include/kmm/runtime/device_resources.hpp +++ /dev/null @@ -1,55 +0,0 @@ -#pragma once - -#include "kmm/core/buffer.hpp" -#include "kmm/core/resource.hpp" -#include "kmm/core/system_info.hpp" -#include "kmm/runtime/stream_manager.hpp" -#include "kmm/utils/gpu_utils.hpp" - -namespace kmm { - -class DeviceResourceOperation { - public: - virtual ~DeviceResourceOperation() = default; - virtual void execute(DeviceResource& resource, std::vector accessors) = 0; -}; - -class DeviceResources { - KMM_NOT_COPYABLE_OR_MOVABLE(DeviceResources) - - public: - DeviceResources( - std::vector contexts, - size_t streams_per_context, - std::shared_ptr stream_manager - ); - - ~DeviceResources(); - - size_t num_contexts() const; - GPUContextHandle context(DeviceId device_id); - - DeviceEvent submit( - DeviceId device_id, - DeviceStreamSet stream_hint, - DeviceEventSet deps, - DeviceResourceOperation& op, - std::vector accessors - ); - - private: - struct Device; - struct Stream; - - Stream* select_stream_for_operation( - DeviceId device_id, - DeviceStreamSet stream_hint, - const DeviceEventSet& deps - ); - - std::shared_ptr m_stream_manager; - size_t m_streams_per_device; - std::vector> m_devices; -}; - -} // namespace kmm diff --git a/include/kmm/runtime/device_stream.hpp b/include/kmm/runtime/device_stream.hpp new file mode 100644 index 00000000..c2f97b9e --- /dev/null +++ b/include/kmm/runtime/device_stream.hpp @@ -0,0 +1,86 @@ +#pragma once + +#include +#include + +#include "fmt/ostream.h" + +#include "kmm/runtime/device_event_registry.hpp" + +namespace kmm { + +class DeviceStream { + KMM_NOT_COPYABLE_OR_MOVABLE(DeviceStream) + + public: + DeviceStream() = default; + + DeviceStream(DeviceEventRegistry registry, DeviceStreamId stream_id) : + m_manager(std::move(registry)), + m_stream_id(stream_id), + m_stream(m_manager.get(stream_id)), + m_context(m_manager.context(stream_id)) {} + + DeviceStreamId id() const noexcept { + return m_stream_id; + } + + g_context_t context() const noexcept { + return m_context; + } + + operator g_stream_t() const noexcept { + return m_stream; + } + + bool is_null() const noexcept { + return m_stream_id.is_null(); + } + + bool preceded_by(const DeviceEventSet& deps) const noexcept { + return m_manager.precedes(deps, m_stream_id); + } + + void wait_on_event(const DeviceEvent& dep) const { + m_manager.wait_on_event(m_stream_id, dep); + } + + void wait_on_event(const DeviceEventSet& deps) const { + m_manager.wait_on_event(m_stream_id, deps); + } + + void wait_on_default_stream() const { + m_manager.wait_on_default_stream(m_stream_id); + } + + void synchronize() const { + m_manager.synchronize(m_stream_id); + } + + DeviceEvent record_event() const { + return m_manager.record(m_stream_id); + } + + /** + * Waits for `deps`, calls `fun(stream)` to submit work onto the stream, then records an event. + */ + template + DeviceEvent submit(const DeviceEventSet& deps, F&& fun) const { + return m_manager.submit(m_stream_id, deps, std::forward(fun)); + } + + friend std::ostream& operator<<(std::ostream& stream, const DeviceStream& e) { + return stream << e.m_stream_id; + } + + private: + DeviceEventRegistry m_manager; + DeviceStreamId m_stream_id; + g_stream_t m_stream = nullptr; + g_context_t m_context = nullptr; +}; + +} // namespace kmm + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; diff --git a/include/kmm/runtime/identifiers.hpp b/include/kmm/runtime/identifiers.hpp new file mode 100644 index 00000000..02941524 --- /dev/null +++ b/include/kmm/runtime/identifiers.hpp @@ -0,0 +1,164 @@ +#pragma once + +#include +#include +#include + +#include "fmt/ostream.h" + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/utils/hash_utils.hpp" + +namespace kmm { + +static constexpr size_t MAX_DEVICES = 4; + +class DeviceId { + public: + KMM_INLINE constexpr explicit DeviceId(size_t id) : m_id(static_cast(id)) { + if (m_id >= MAX_DEVICES) { + throw std::runtime_error("Device id out of range"); + } + } + + KMM_INLINE size_t get() const noexcept { + KMM_UNSAFE_ASSUME(m_id < MAX_DEVICES); + return m_id; + } + + KMM_INLINE constexpr bool operator==(const DeviceId& that) const noexcept { + return m_id == that.m_id; + } + + KMM_INLINE constexpr bool operator!=(const DeviceId& that) const noexcept { + return !(*this == that); + } + + friend std::ostream& operator<<(std::ostream& stream, const DeviceId& e); + + private: + uint8_t m_id; +}; + +class BufferId { + public: + KMM_INLINE constexpr explicit BufferId(uint64_t id) : m_id(id) {} + + KMM_INLINE constexpr uint64_t get() const noexcept { + return m_id; + } + + KMM_INLINE constexpr bool operator==(const BufferId& that) const noexcept { + return m_id == that.m_id; + } + + KMM_INLINE constexpr bool operator!=(const BufferId& that) const noexcept { + return !(*this == that); + } + + friend std::ostream& operator<<(std::ostream& stream, const BufferId& e); + + private: + uint64_t m_id; +}; + +class MemoryId { + public: + enum struct Kind { Host, Device }; + + KMM_INLINE static constexpr MemoryId host() { + return MemoryId {Kind::Host, DeviceId(0)}; + } + + KMM_INLINE static constexpr MemoryId device(DeviceId id) { + return MemoryId {Kind::Device, id}; + } + + // Parses strings such as "host", "cpu", "gpu", "cuda", "gpu:0", "cuda:1", "device:2". + MemoryId(const std::string& name); + + KMM_INLINE + bool is_host() const noexcept { + return m_kind == Kind::Host; + } + + KMM_INLINE + bool is_device() const noexcept { + return !is_host(); + } + + KMM_INLINE + DeviceId as_device() const noexcept { + KMM_ASSERT(is_device()); + return m_device_id; + } + + KMM_INLINE constexpr bool operator==(const MemoryId& that) const noexcept { + return m_kind == that.m_kind && (m_kind != Kind::Device || m_device_id == that.m_device_id); + } + + KMM_INLINE constexpr bool operator!=(const MemoryId& that) const noexcept { + return !(*this == that); + } + + KMM_INLINE constexpr bool operator<(const MemoryId& that) const noexcept { + if (m_kind == Kind::Device && that.m_kind == Kind::Device) { + return m_device_id.get() < that.m_device_id.get(); + } else { + return m_kind < that.m_kind; + } + } + + KMM_INLINE constexpr bool operator>(const MemoryId& that) const noexcept { + return that < *this; + } + + KMM_INLINE constexpr bool operator<=(const MemoryId& that) const noexcept { + return *this < that || *this == that; + } + + KMM_INLINE constexpr bool operator>=(const MemoryId& that) const noexcept { + return that <= *this; + } + + friend std::ostream& operator<<(std::ostream& stream, const MemoryId& e); + + private: + constexpr MemoryId(Kind kind, DeviceId device_id) : m_kind(kind), m_device_id(device_id) {} + + Kind m_kind; + DeviceId m_device_id; +}; + +} // namespace kmm + +template<> +struct std::hash { + size_t operator()(kmm::BufferId val) const noexcept { + return ::kmm::hash_fields(val.get()); + } +}; + +template<> +struct std::hash { + size_t operator()(kmm::DeviceId val) const noexcept { + return ::kmm::hash_fields(val.get()); + } +}; + +template<> +struct std::hash { + size_t operator()(kmm::MemoryId val) const noexcept { + return val.is_device() ? val.as_device().get() : size_t(-1); + } +}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/runtime/memops/copy.hpp b/include/kmm/runtime/memops/copy.hpp new file mode 100644 index 00000000..7b12ead6 --- /dev/null +++ b/include/kmm/runtime/memops/copy.hpp @@ -0,0 +1,150 @@ +#pragma once + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/checked_math.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/core/range.hpp" +#include "kmm/runtime/memops/types.hpp" + +namespace kmm { + +/// \addtogroup memops +/// @{ + +/// Describes a single axis of a strided copy: `extent` elements are copied, and each successive +/// element is offset by `src_stride`/`dst_stride` bytes in the source/destination buffer +/// respectively. +struct CopyDim { + memops_extent_type extent = 1; + memops_stride_type src_stride = 0; + memops_stride_type dst_stride = 0; +}; + +/// Describes a (possibly strided, possibly multi-dimensional) copy from a source buffer to a +/// destination buffer. +struct CopyDescription { + /// The size (in bytes) of a single element. + size_t element_size = 1; + + /// A byte offset added to `src_addr` before applying `dims`. + memops_stride_type src_offset = 0; + + /// A byte offset added to `dst_addr` before applying `dims`. + memops_stride_type dst_offset = 0; + + /// The number of axes described by `dims`. Must be at most `MEMOPS_MAX_DIMS`. + size_t num_dims = 0; + + /// The extent and per-axis strides, ordered from the outermost to the innermost axis. + CopyDim dims[MEMOPS_MAX_DIMS] = {}; + + CopyDescription() = default; + + KMM_HOST_DEVICE + explicit CopyDescription(size_t element_size) : element_size(element_size) {} + + /// Appends an axis to this description. `src_stride`/`dst_stride` must be given in bytes. + /// + /// An axis of extent one is dropped (it is visited exactly once, so its stride never + /// contributes to addressing). + KMM_HOST_DEVICE + void add_dimension( + memops_extent_type extent, + memops_stride_type src_stride, + memops_stride_type dst_stride + ) { + if (extent == 1) { + return; + } + + KMM_ASSERT(num_dims < MEMOPS_MAX_DIMS); + dims[num_dims] = CopyDim {extent, src_stride, dst_stride}; + num_dims++; + } + + /// Returns the total number of elements copied (the product of the extent of each axis). + KMM_HOST_DEVICE + memops_extent_type num_elements() const { + memops_extent_type result = 1; + + for (size_t i = 0; i < num_dims; i++) { + result *= dims[i].extent; + } + + return result; + } + + /// Returns `true` if this copy reads and writes nothing because some axis has extent zero. + KMM_HOST_DEVICE + bool is_empty() const { + return num_elements() == 0; + } + + /// Returns the half-open range of byte offsets (relative to `src_addr`) that this copy + /// will read from. + Range src_range() const; + + /// Returns the half-open range of byte offsets (relative to `dst_addr`) that this copy + /// will write to. + Range dst_range() const; + + /// Returns an equivalent description in canonical form: axes are sorted by decreasing stride + /// (largest first) and adjacent contiguous axes are merged, so backends (e.g. `copy_gpu`) that + /// only special-case a few axes see the smallest `num_dims`. An empty copy (some axis of + /// extent zero) collapses to a single axis with `extent == 0` and all strides zero; see + /// `is_empty`. + CopyDescription simplify() const; +}; + +/// Builds a `CopyDescription` that copies `element_size`-sized elements between two layouts of +/// matching rank (e.g. `kmm::Layout`), translating each layout's (element-space) base offset, +/// per-axis origin, and strides into the byte offsets/strides that `CopyDescription` expects. +template +CopyDescription make_copy_description( + const DstLayoutT& dst, + const SrcLayoutT& src, + size_t element_size +) { + static_assert(DstLayoutT::rank == SrcLayoutT::rank, "rank mismatch"); + static_assert(DstLayoutT::rank <= MEMOPS_MAX_DIMS, "rank exceeds maximum"); + + ptrdiff_t dst_offset = dst.base_offset(); + ptrdiff_t src_offset = src.base_offset(); + + for (size_t i = 0; i < DstLayoutT::rank; i++) { + dst_offset += static_cast(dst.stride(i)) * static_cast(dst.begin(i)); + src_offset += static_cast(src.stride(i)) * static_cast(src.begin(i)); + } + + CopyDescription descr(element_size); + descr.src_offset = checked_mul(src_offset, element_size); + descr.dst_offset = checked_mul(dst_offset, element_size); + + for (size_t i = 0; i < DstLayoutT::rank; i++) { + descr.add_dimension( + checked_cast(dst.extent(i)), + checked_mul(src.stride(i), element_size), + checked_mul(dst.stride(i), element_size) + ); + } + + return descr.simplify(); +} + +/// @} + +namespace memops { + +/// \addtogroup memops +/// @{ + +/// Copies data from `src_addr` to `dst_addr` on the CPU, according to `description`. Blocks +/// until the copy has completed. +void copy(const void* src_addr, void* dst_addr, const CopyDescription& description); + +/// @} + +} // namespace memops + +} // namespace kmm diff --git a/include/kmm/runtime/memops/copy_gpu.hpp b/include/kmm/runtime/memops/copy_gpu.hpp new file mode 100644 index 00000000..82053d86 --- /dev/null +++ b/include/kmm/runtime/memops/copy_gpu.hpp @@ -0,0 +1,25 @@ +#pragma once + +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/device_stream.hpp" +#include "kmm/runtime/memops/copy.hpp" + +namespace kmm::memops { + +/// \addtogroup memops +/// @{ + +/// Copies data from `src_addr` to `dst_addr` on the GPU, according to `description`. The copy is +/// enqueued on `stream` after waiting for `dependencies`, and the returned event becomes ready +/// once the copy has completed. Both `src_addr` and `dst_addr` must point to device memory on +/// the same device (i.e. this is a device-to-device copy). +void copy_gpu( + g_stream_t stream, + const void* src_addr, + void* dst_addr, + const CopyDescription& description +); + +/// @} + +} // namespace kmm::memops diff --git a/include/kmm/runtime/memops/fill.hpp b/include/kmm/runtime/memops/fill.hpp new file mode 100644 index 00000000..c558e708 --- /dev/null +++ b/include/kmm/runtime/memops/fill.hpp @@ -0,0 +1,108 @@ +#pragma once + +#include +#include +#include + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/types.hpp" + +namespace kmm { + +/// \addtogroup memops +/// @{ + +/// The raw bytes of some POD type. For example. use `FillValue::from(int(5))` to construct a `FillValue`. +struct FillValue { + static constexpr size_t MAX_FILL_LENGTH = 32; + + FillValue() = default; + + template + static FillValue from(T value) { + static_assert(sizeof(T) <= MAX_FILL_LENGTH, "T exceeds capacity of FillValue"); + static_assert( + std::is_trivially_copyable_v, + "T must be trivially copyable to be stored in a FillValue" + ); + + FillValue result; + result.length = sizeof(T); + std::memcpy(result.buffer, &value, sizeof(T)); + return result; + } + + size_t length = 0; + std::byte buffer[MAX_FILL_LENGTH] {}; +}; + +/// Describes a single axis of a strided fill. +struct FillDim { + memops_extent_type extent = 1; // number of elements + memops_stride_type stride = 0; // step size between elements in bytes. +}; + +/// Describes a strided multi-dimensional fill of a destination buffer with a repeating element value. +/// +/// The destination is treated as an opaque byte buffer: `fill`/`fill_gpu` do not interpret the +/// bytes being written, so `element_size` only determines the width of the value written at each +/// position (see `fill`/`fill_gpu`). +struct FillDescription { + FillValue value; + + /// Offset add to `dst_addr`, in bytes. + memops_stride_type offset = 0; + + /// The number of axes described by `dims`. Must be at most `MEMOPS_MAX_DIMS`. + size_t num_dims = 0; + + /// The extent and per-axis stride. + FillDim dims[MEMOPS_MAX_DIMS] = {}; + + FillDescription() = default; + + explicit FillDescription(FillValue value) : value(value) {} + + /// Appends an axis to this description. `stride` must be given in bytes. + void add_dimension(memops_extent_type extent, memops_stride_type stride); + + /// Returns an equivalent description with `dims` sorted from the largest stride to the + /// smallest, and adjacent axes merged whenever they are contiguous (i.e. the outer axis's + /// stride equals the inner axis's stride times its extent). This reduces the number of + /// dimensions and may make certain operations easier. + FillDescription simplify() const; + + /// Returns the total number of elements written (the product of the extent of each axis). + memops_extent_type num_elements() const { + memops_extent_type result = 1; + + for (size_t i = 0; i < num_dims; i++) { + result *= dims[i].extent > 0 ? dims[i].extent : memops_extent_type {0}; + } + + return result; + } + + /// Returns `true` if this fill writes nothing because some axis has extent zero. + bool is_empty() const { + return num_elements() == 0; + } +}; + +/// @} + +namespace memops { + +/// \addtogroup memops +/// @{ + +/// Fills `dst_addr` on the CPU, according to `description`, with copies of the `element_size` +/// bytes pointed to by `fill_value`. Blocks until the fill has completed. +void fill(void* dst_addr, const FillDescription& description); + +/// @} + +} // namespace memops + +} // namespace kmm diff --git a/include/kmm/runtime/memops/fill_gpu.hpp b/include/kmm/runtime/memops/fill_gpu.hpp new file mode 100644 index 00000000..5fd9dd02 --- /dev/null +++ b/include/kmm/runtime/memops/fill_gpu.hpp @@ -0,0 +1,18 @@ +#pragma once + +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/utils/backends.hpp" + +namespace kmm::memops { + +/// \addtogroup memops +/// @{ + +/// Fills `dst_addr` on the GPU, according to `description`, with copies of the fill value held in +/// `description.value` (the value bytes are read on the host before the work is enqueued). The +/// fill is enqueued asynchronously on `stream`. +void fill_gpu(g_stream_t stream, void* dst_addr, const FillDescription& description); + +/// @} + +} // namespace kmm::memops diff --git a/include/kmm/runtime/memops/reducer.hpp b/include/kmm/runtime/memops/reducer.hpp new file mode 100644 index 00000000..e63bd409 --- /dev/null +++ b/include/kmm/runtime/memops/reducer.hpp @@ -0,0 +1,155 @@ +#pragma once + +#include +#include + +#include "kmm/core/key_value.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/runtime/memops/types.hpp" + +namespace kmm::memops { + +/// Per-`ReductionOp` accumulator +template +struct Reducer; + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value += x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + T value = 0; +}; + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value *= x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + T value = 1; +}; + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value = value < x ? value : x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + static constexpr T identity_value = std::numeric_limits::max(); + T value = identity_value; +}; + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value = value > x ? value : x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + static constexpr T identity_value = std::numeric_limits::lowest(); + T value = identity_value; +}; + +// `BitwiseAnd`/`BitwiseOr` are defined for integer `T` only. + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value &= x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + T value = static_cast(~static_cast(0)); +}; + +template +struct Reducer>> { + using element_type = T; + + KMM_HOST_DEVICE void consume(T x) { + value |= x; + } + + KMM_HOST_DEVICE T finish() { + return value; + } + + T value = static_cast(0); +}; + +// `KeyValue` (argmin/argmax): ordered by value with the key as tie-breaker, so only `Min`/`Max` +// are defined. The identity carries the neutral value (the largest/smallest `VT`) and key `0`. + +template +struct Reducer, ReductionOp::Min> { + using element_type = KeyValue; + + KMM_HOST_DEVICE void consume(KeyValue x) { + value = value < x ? value : x; + } + + KMM_HOST_DEVICE KeyValue finish() { + return value; + } + + static constexpr VT identity_value = std::numeric_limits::max(); + KeyValue value = {0, identity_value}; +}; + +template +struct Reducer, ReductionOp::Max> { + using element_type = KeyValue; + + KMM_HOST_DEVICE void consume(KeyValue x) { + value = value > x ? value : x; + } + + KMM_HOST_DEVICE KeyValue finish() { + return value; + } + + static constexpr VT identity_value = std::numeric_limits::lowest(); + KeyValue value = {0, identity_value}; +}; + +/// `true` when `Reducer` is a complete type, i.e. `Op` is a supported reduction for element +/// type `T`. Used by the dispatch to reject unsupported combinations (e.g. `BitwiseAnd` on `float`, +/// `Sum` on `KeyValue`) with a runtime error instead of a compile error. +template +inline constexpr bool is_reduction_supported = false; + +template +inline constexpr bool is_reduction_supported))>> = + true; + +} // namespace kmm::memops diff --git a/include/kmm/runtime/memops/reduction.hpp b/include/kmm/runtime/memops/reduction.hpp new file mode 100644 index 00000000..e5137150 --- /dev/null +++ b/include/kmm/runtime/memops/reduction.hpp @@ -0,0 +1,383 @@ +#pragma once + +#include +#include + +#include "kmm/core/checked_math.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/copy.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/types.hpp" + +namespace kmm { + +/// The identity element for `op` at `dtype`, as the raw bytes to broadcast-fill a buffer that is +/// about to be reduced into (so slots never written still drop out of the fold): `0` for `Sum`, +/// `1` for `Product`, the type's maximum for `Min`, its lowest for `Max`. +FillValue reduction_identity(DataType dtype, ReductionOp op); + +/// \addtogroup memops +/// @{ + +/// Describes a single "batch" axis of a strided reduction: `extent` output elements are +/// produced, and each successive output element is offset by `input_stride`/`output_stride` +/// bytes in the input/output buffer respectively. +/// +/// This describes the axes that are *not* reduced over; see `ReductionDescription::reduction_extent` +/// for the axis that is reduced over. +struct ReductionDim { + memops_extent_type extent = 1; + memops_stride_type input_stride = 0; + memops_stride_type output_stride = 0; +}; + +/// Describes a (possibly strided, possibly multi-dimensional) reduction of an input buffer into +/// an output buffer. +/// +/// A reduction combines groups of input elements into a single output element using `operation` +/// (e.g. `Sum`). `dims` describes the "batch" axes: one output element is produced per +/// combination of batch indices. `reduction_extent`/`reduction_stride` describe the additional +/// axis that is reduced over: for each output element, `reduction_extent` input elements +/// (spaced `reduction_stride` bytes apart, starting at the corresponding batch position) are +/// combined together. +/// +/// For example, reducing an `(m, n)` row-major input down to `m` outputs (i.e. reducing over the +/// trailing axis) is described by a single batch dimension with `extent = m` and +/// `input_stride = n * element_size`, together with `reduction_extent = n` and +/// `reduction_stride = element_size`. +struct ReductionDescription { + /// The element type operated on. Determines both how elements are interpreted numerically + /// and the size (in bytes) of a single element. + DataType dtype = DataType::Float32; + + /// The operator used to combine elements. + ReductionOp operation = ReductionOp::Sum; + + /// A byte offset added to `src_addr` before applying `dims`/`reduction_stride`. + memops_stride_type input_offset = 0; + + /// A byte offset added to `dst_addr` before applying `dims`. + memops_stride_type output_offset = 0; + + /// The number of batch axes described by `dims`. Must be at most `MEMOPS_MAX_DIMS`. + size_t num_dims = 0; + + /// The extent and per-axis strides of the batch axes, ordered from the outermost to the + /// innermost axis. + ReductionDim dims[MEMOPS_MAX_DIMS] = {}; + + /// The number of input elements combined into each output element. + memops_extent_type reduction_extent = 1; + + /// The byte offset between successive input elements being combined into the same output + /// element. + memops_stride_type reduction_stride = 0; + + /// If `true`, each output element is combined with its previous value (i.e. the output + /// buffer is treated as already holding a partial result). If `false`, each output element + /// is overwritten with the result of the reduction. + bool accumulate = false; + + ReductionDescription() = default; + + KMM_HOST_DEVICE + ReductionDescription(DataType dtype, ReductionOp operation) : + dtype(dtype), + operation(operation) {} + + /// Appends a batch axis to this description. `input_stride`/`output_stride` must be given in + /// bytes. + /// + /// An axis of extent one is dropped (it is visited exactly once, so its stride never + /// contributes to addressing). Otherwise, this axis is checked against every existing axis + /// for contiguity (in both `input_stride` and `output_stride`) and folded into the first one + /// it is contiguous with instead of consuming a new slot. This keeps `num_dims` from growing + /// unnecessarily, which matters since `dims` only has room for `MEMOPS_MAX_DIMS` axes. + KMM_HOST_DEVICE + void add_dimension( + memops_extent_type extent, + memops_stride_type input_stride, + memops_stride_type output_stride + ) { + if (extent == 1) { + return; + } + + for (size_t i = 0; i < num_dims; i++) { + ReductionDim& dim = dims[i]; + + // `dim` is the outer neighbor of the new axis: it keeps its extent, but adopts the + // new axis's (smaller) stride as its own. + if (dim.input_stride == input_stride * extent + && dim.output_stride == output_stride * extent) { + dim.extent *= extent; + dim.input_stride = input_stride; + dim.output_stride = output_stride; + return; + } + + // `dim` is the inner neighbor of the new axis: the new axis's stride already + // matches `dim`'s span, so only `dim`'s extent needs to grow. + if (input_stride == dim.input_stride * dim.extent + && output_stride == dim.output_stride * dim.extent) { + dim.extent *= extent; + return; + } + } + + KMM_ASSERT(num_dims < MEMOPS_MAX_DIMS); + dims[num_dims] = ReductionDim {extent, input_stride, output_stride}; + num_dims++; + } + + /// Returns the total number of output elements produced (the product of the extent of each + /// batch axis). + KMM_HOST_DEVICE + memops_extent_type num_outputs() const { + memops_extent_type result = 1; + + for (size_t i = 0; i < num_dims; i++) { + result *= dims[i].extent; + } + + return result; + } + + /// Returns the half-open range of byte offsets (relative to `src_addr`) that this reduction + /// reads from. Accounts for `input_offset`, the batch axes (`dims`), and the reduced axis + /// (`reduction_extent`/`reduction_stride`). + Range src_range() const; + + /// Returns the half-open range of byte offsets (relative to `dst_addr`) that this reduction + /// writes to. Accounts for `output_offset` and the batch axes (`dims`). + Range dst_range() const; + + /// Returns an equivalent description with adjacent contiguous batch axes (`dims`) merged, to + /// keep `num_dims` small. + ReductionDescription simplify() const; + + /// Returns `true` if this reduction combines exactly one input element into each output + /// element without accumulating into the previous value: every `combine(identity, value)` + /// then reduces to `value` (for all of `Sum`/`Product`/`Min`/`Max`/`BitwiseAnd`/`BitwiseOr`), + /// so the reduction is equivalent to a plain strided copy from `src_addr` to `dst_addr`. See + /// `as_copy`. + KMM_HOST_DEVICE + bool is_equivalent_to_copy() const { + return reduction_extent == 1 && !accumulate; + } + + /// Converts this reduction into an equivalent `CopyDescription`. Only valid when + /// `is_equivalent_to_copy()` returns `true`. + CopyDescription as_copy() const; + + /// Returns `true` if this reduction combines *zero* input elements into each output element + /// (`reduction_extent == 0`) without accumulating into the previous value: every output is + /// therefore left as the identity element for `operation`, so the reduction is equivalent to + /// broadcasting that identity across the output region. See `as_fill`. + KMM_HOST_DEVICE + bool is_equivalent_to_fill() const { + return reduction_extent == 0 && !accumulate; + } + + /// Converts this reduction into an equivalent `FillDescription` that writes the identity + /// element of `operation` across the output region (the batch axes). Only valid when + /// `is_equivalent_to_fill()` returns `true`. + FillDescription as_fill() const; + + /// Returns `true` if this reduction is guaranteed to leave the output buffer unchanged, so it + /// need not read or write anything. That happens when no output elements are produced at all + /// (`num_outputs() == 0`), or when zero input elements are combined into each output + /// (`reduction_extent == 0`) while accumulating: every output becomes + /// `combine(previous, identity) == previous` for all of + /// `Sum`/`Product`/`Min`/`Max`/`BitwiseAnd`/`BitwiseOr`. + KMM_HOST_DEVICE + bool is_noop() const { + return num_outputs() == 0 || (reduction_extent == 0 && accumulate); + } +}; + +/// Builds a `ReductionDescription` that reduces `src` (e.g. a `kmm::Layout`) over `axis` into +/// `dst`, whose rank must be exactly one less than `src`'s: every axis of `src` other than `axis` +/// maps (in order) to the corresponding axis of `dst`, and `axis` itself becomes the reduced +/// axis. Translates each layout's (element-space) base offset, per-axis origin, and strides into +/// the byte offsets/strides that `ReductionDescription` expects. +template +ReductionDescription make_reduction_description( + const DstLayoutT& dst, + const SrcLayoutT& src, + size_t axis, + DataType dtype, + ReductionOp op +) { + static_assert( + DstLayoutT::rank + 1 == SrcLayoutT::rank, + "dst must have exactly one axis fewer than src" + ); + static_assert(SrcLayoutT::rank <= MEMOPS_MAX_DIMS + 1, "rank exceeds maximum"); + KMM_ASSERT(axis < SrcLayoutT::rank); + + size_t element_size = data_type_size(dtype); + + ptrdiff_t dst_offset = dst.base_offset(); + ptrdiff_t src_offset = src.base_offset(); + + for (size_t i = 0; i < DstLayoutT::rank; i++) { + dst_offset += static_cast(dst.stride(i)) * static_cast(dst.begin(i)); + } + + for (size_t i = 0; i < SrcLayoutT::rank; i++) { + src_offset += static_cast(src.stride(i)) * static_cast(src.begin(i)); + } + + ReductionDescription descr(dtype, op); + descr.input_offset = checked_mul(src_offset, element_size); + descr.output_offset = checked_mul(dst_offset, element_size); + descr.reduction_extent = checked_cast(src.extent(axis)); + descr.reduction_stride = checked_mul(src.stride(axis), element_size); + + size_t dst_axis = 0; + + for (size_t i = 0; i < SrcLayoutT::rank; i++) { + if (i == axis) { + continue; + } + + descr.add_dimension( + checked_cast(dst.extent(dst_axis)), + checked_mul(src.stride(i), element_size), + checked_mul(dst.stride(dst_axis), element_size) + ); + + dst_axis++; + } + + return descr; +} + +/// Multi-axis counterpart of `make_reduction_description`: reduces `src` over *every* axis listed +/// in `axes` at once. This collapses into a single `ReductionDescription` only when those axes +/// occupy one contiguous run of memory -- sorted by stride, each axis starts exactly where the +/// previous one ends -- so they fold into one reduced axis. Returns `std::nullopt` when they do +/// not, leaving the caller to reduce one axis at a time. +/// +/// `dst`'s rank must be `src`'s rank minus `axes.size()`; every axis of `src` not in `axes` maps, +/// in order, onto the corresponding axis of `dst`. Reduced axes of extent one are ignored (they +/// never contribute to the fold). +template +std::optional make_reduction_description( + const DstLayoutT& dst, + const SrcLayoutT& src, + std::initializer_list axes, + DataType dtype, + ReductionOp op +) { + static_assert(SrcLayoutT::rank >= 1, "src must have at least one axis"); + static_assert(SrcLayoutT::rank <= MEMOPS_MAX_DIMS + 1, "rank exceeds maximum"); + KMM_ASSERT(DstLayoutT::rank + axes.size() == SrcLayoutT::rank); + + size_t element_size = data_type_size(dtype); + + bool is_reduced[SrcLayoutT::rank] = {}; + + for (size_t axis : axes) { + KMM_ASSERT(axis < SrcLayoutT::rank); + KMM_ASSERT(!is_reduced[axis]); + is_reduced[axis] = true; + } + + // Collect the reduced axes in element space, dropping extent-one axes, then sort them by + // ascending stride so a contiguous run reads as `stride[i] == stride[i - 1] * extent[i - 1]`. + ptrdiff_t run_extent[SrcLayoutT::rank]; + ptrdiff_t run_stride[SrcLayoutT::rank]; + size_t run_count = 0; + + for (size_t i = 0; i < SrcLayoutT::rank; i++) { + if (is_reduced[i] && src.extent(i) != 1) { + run_extent[run_count] = static_cast(src.extent(i)); + run_stride[run_count] = static_cast(src.stride(i)); + run_count++; + } + } + + for (size_t i = 1; i < run_count; i++) { + for (size_t j = i; j > 0 && run_stride[j] < run_stride[j - 1]; j--) { + ptrdiff_t tmp_extent = run_extent[j]; + run_extent[j] = run_extent[j - 1]; + run_extent[j - 1] = tmp_extent; + + ptrdiff_t tmp_stride = run_stride[j]; + run_stride[j] = run_stride[j - 1]; + run_stride[j - 1] = tmp_stride; + } + } + + ptrdiff_t merged_extent = 1; + ptrdiff_t merged_stride = 0; + + for (size_t i = 0; i < run_count; i++) { + if (i == 0) { + merged_stride = run_stride[0]; + merged_extent = run_extent[0]; + } else if (run_stride[i] == merged_stride * merged_extent) { + merged_extent *= run_extent[i]; + } else { + return std::nullopt; + } + } + + ptrdiff_t dst_offset = dst.base_offset(); + ptrdiff_t src_offset = src.base_offset(); + + for (size_t i = 0; i < DstLayoutT::rank; i++) { + dst_offset += static_cast(dst.stride(i)) * static_cast(dst.begin(i)); + } + + for (size_t i = 0; i < SrcLayoutT::rank; i++) { + src_offset += static_cast(src.stride(i)) * static_cast(src.begin(i)); + } + + ReductionDescription descr(dtype, op); + descr.input_offset = checked_mul(src_offset, element_size); + descr.output_offset = checked_mul(dst_offset, element_size); + descr.reduction_extent = checked_cast(merged_extent); + descr.reduction_stride = checked_mul(merged_stride, element_size); + + size_t dst_axis = 0; + + for (size_t i = 0; i < SrcLayoutT::rank; i++) { + if (is_reduced[i]) { + continue; + } + + descr.add_dimension( + checked_cast(dst.extent(dst_axis)), + checked_mul(src.stride(i), element_size), + checked_mul(dst.stride(dst_axis), element_size) + ); + + dst_axis++; + } + + return descr; +} + +/// @} + +namespace memops { + +/// \addtogroup memops +/// @{ + +/// Reduces `src_addr` into `dst_addr` on the CPU, according to `description`. Blocks until the +/// reduction has completed. +/// +/// `src_addr`/`dst_addr` and every offset/stride in `description` are assumed to be aligned to the +/// element type; descriptions built by `make_reduction_description` always satisfy this. +void reduce(const void* src_addr, void* dst_addr, const ReductionDescription& description); + +/// @} + +} // namespace memops + +} // namespace kmm diff --git a/include/kmm/runtime/memops/reduction_gpu.hpp b/include/kmm/runtime/memops/reduction_gpu.hpp new file mode 100644 index 00000000..fdd4dd9b --- /dev/null +++ b/include/kmm/runtime/memops/reduction_gpu.hpp @@ -0,0 +1,28 @@ +#pragma once + +#include "kmm/runtime/memops/copy_gpu.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/utils/backends.hpp" + +namespace kmm::memops { + +/// \addtogroup memops +/// @{ + +/// Reduces `src_addr` into `dst_addr` on the GPU, according to `description`. The reduction is +/// enqueued on `stream` after waiting for `dependencies`, and the returned event becomes ready +/// once the reduction has completed. +void reduce_gpu( + g_stream_t stream, + const void* src_addr, + void* dst_addr, + void* scratch_addr, + const ReductionDescription& description +); + +/// Returns the size in bytes of the scratch buffer that `reduce_gpu` requires for `description`. +size_t reduce_gpu_scratch_size(const ReductionDescription& description); + +/// @} + +} // namespace kmm::memops diff --git a/include/kmm/runtime/memops/types.hpp b/include/kmm/runtime/memops/types.hpp new file mode 100644 index 00000000..40646a44 --- /dev/null +++ b/include/kmm/runtime/memops/types.hpp @@ -0,0 +1,104 @@ +#pragma once + +#include + +#include "kmm/core/key_value.hpp" +#include "kmm/core/macros.hpp" + +namespace kmm { + +/// \addtogroup memops +/// @{ + +/// The maximum number of dimensions supported by the strided descriptors in `kmm/memops` +/// (`CopyDescription`, `FillDescription`, `ReductionDescription`). Kept small and fixed so these +/// descriptors stay plain, fixed-size, trivially-copyable structs that can be passed by value to +/// GPU kernels. +inline constexpr size_t MEMOPS_MAX_DIMS = 4; + +/// The extent (number of elements) along a single axis of a `kmm/memops` descriptor. +using memops_extent_type = signed long long int; + +/// A byte offset between consecutive elements along a single axis of a `kmm/memops` descriptor. +using memops_stride_type = signed long long int; + +/// A runtime tag for the scalar element type operated on by `reduce`. +enum class DataType { + Unknown = 0, + Int32, + Int64, + Uint32, + Uint64, + Float32, + Float64, + /// `KeyValue` / `KeyValue`: a value paired with its `int64_t` key. Only + /// meaningful for `Min`/`Max` reductions (i.e. argmin/argmax). + KeyValueInt64, + KeyValueFloat64, +}; + +/// Returns the size (in bytes) of a single element of the given data type. +size_t data_type_size(DataType dtype); + +/// Returns a human-readable name for the given data type (e.g. `"Float32"`). +const char* data_type_name(DataType dtype); + +/// The operator applied by `reduce`/`reduce_gpu` to combine elements. +enum class ReductionOp { + Sum, + Product, + Min, + Max, + BitwiseAnd, + BitwiseOr, +}; + +/// Returns a human-readable name for the given reduction operator (e.g. `"Sum"`). +const char* reduction_op_name(ReductionOp op); + +/// Properties of a `DataType` tag, keyed on the tag itself. Each specialization provides: +/// - `element_type`: the C++ type stored in each element; +/// - `name`: a human-readable name (e.g. `"Float32"`). +template +struct data_type_traits; + +/// The C++ element type corresponding to the `DataType` tag `dtype`. +template +using element_type_t = typename data_type_traits::element_type; + +/// Maps a C++ element type to its `DataType` tag via a `static constexpr DataType value` member. +template +struct data_type_of_impl; + +/// The `DataType` tag for `T`. Valid for every built-in element type and any type with a +/// `data_type_of_impl` specialization. +template +constexpr DataType data_type_of() { + return data_type_of_impl::value; +} + +#define KMM_IMPL_DATA_TYPE(TYPE, DTYPE) \ + template<> \ + struct data_type_traits { \ + using element_type = TYPE; \ + static constexpr const char* name = #DTYPE; \ + }; \ + template<> \ + struct data_type_of_impl { \ + static constexpr DataType value = DataType::DTYPE; \ + }; + +KMM_IMPL_DATA_TYPE(int32_t, Int32) +KMM_IMPL_DATA_TYPE(int64_t, Int64) +KMM_IMPL_DATA_TYPE(uint32_t, Uint32) +KMM_IMPL_DATA_TYPE(uint64_t, Uint64) +KMM_IMPL_DATA_TYPE(float, Float32) +KMM_IMPL_DATA_TYPE(double, Float64) +KMM_IMPL_DATA_TYPE(KeyValue, KeyValueInt64) +KMM_IMPL_DATA_TYPE(KeyValue, KeyValueFloat64) + +#undef KMM_IMPL_DATA_TYPE + +/// @} + +} // namespace kmm diff --git a/include/kmm/runtime/memory_buffer.hpp b/include/kmm/runtime/memory_buffer.hpp new file mode 100644 index 00000000..26a661d5 --- /dev/null +++ b/include/kmm/runtime/memory_buffer.hpp @@ -0,0 +1,369 @@ +#pragma once + +#include + +#include "kmm/runtime/data_interfaces/base.hpp" +#include "kmm/runtime/memory_manager.hpp" +#include "kmm/utils/refcnt_ptr.hpp" + +namespace kmm { + +struct BufferQueueNode { + BufferQueueNode(MemoryId memory_id, AccessKind mode, NotifyHandle callback) : + memory_id(memory_id), + mode(mode) {} + + // Intrusive doubly-linked list used by `MemoryBufferImpl`'s granted/pending queue. + BufferQueueNode* queue_prev = nullptr; + BufferQueueNode* queue_next = nullptr; + MemoryId memory_id; + AccessKind mode; +}; + +struct AccessControl { + // epoch_events: events created by Epoch access. All future requests must wait for + // all epoch events to finish before they can access the data. + // write_events: events created by Epoch or Write access. Future reads must + // wait for all write events before they may safely observe the data. + // read_events: events created by Epoch or Write or Read access. Future writes + // must wait for all read events before they may safely overwrite the data. + // + // The ordering is always: + // - if event in epoch_events => event in write_events + // - if event in write_events => event in read_events + DeviceEventSet epoch_events; + DeviceEventSet write_events; + DeviceEventSet read_events; + + // Whether this location holds (or will hold, once any producer finishes) semantically + // correct data. For `HostAccessControl`, this is independent of `pending_future`: a + // location can be `is_valid == true` while a fill is still in flight (readers must drain + // `pending_future` first), and it can also be `is_valid == false` while a *stale* future + // from before an invalidation is still draining in the background (see + // `MemoryBufferImpl::invalidate_other_allocs`/`invalidate_all`). + bool is_valid = false; + bool is_allocated = false; + size_t alloc_count = 0; + + // Records an access of the given `mode`, maintaining the epoch/write/read + // nesting invariant documented above. + void record_access(AccessKind mode, const DeviceEventSet& deps) noexcept { + if (mode == AccessKind::Exclusive) { + epoch_events.insert(deps); + } + + if (mode != AccessKind::ReadOnly) { + write_events.insert(deps); + } + + read_events.insert(deps); + } + + const DeviceEventSet& retrieve_access(AccessKind mode) noexcept { + switch (mode) { + case AccessKind::ReadOnly: + return read_events; + case AccessKind::SharedWrite: + return write_events; + case AccessKind::Exclusive: + return epoch_events; + default: + KMM_PANIC("invalid state"); + } + } + + void mark_allocated(DeviceEventSet deps) { + KMM_ASSERT(alloc_count == 0); + KMM_ASSERT(!is_allocated); + is_allocated = true; + is_valid = false; + write_events = deps; + read_events = deps; + epoch_events = std::move(deps); + } + + void mark_valid(const DeviceEventSet& events) noexcept { + epoch_events.insert(events); + write_events.insert(events); + read_events.insert(events); + is_valid = true; + } + + void mark_allocated_and_valid(const DeviceEventSet& events) noexcept { + KMM_ASSERT(alloc_count == 0); + KMM_ASSERT(!is_allocated); + mark_valid(events); + is_allocated = true; + } + + // Returns the dependencies the caller must wait on before actually freeing the + // underlying allocation (via `DataInterface`), then clears all state. + DeviceEventSet mark_deallocated() noexcept { + DeviceEventSet deps = retrieve_access(AccessKind::ReadOnly); + KMM_ASSERT(alloc_count == 0); + is_allocated = false; + is_valid = false; + epoch_events.clear(); + write_events.clear(); + read_events.clear(); + return deps; + } +}; + +struct HostAccessControl: AccessControl { + // Tracks an in-flight async host producer (currently only `DataInterface::initialize_host`). + // Mostly independent of `is_valid`: the caller decides what a drained future implies (e.g. + // `ensure_alloc_valid` marks the location valid once it starts a fresh producer; a stale + // future left over from an invalidation is just discarded once drained). The one exception + // is failure: if the producer throws, `is_valid` is forced back to `false` here, since a + // failed producer never actually produced valid data, regardless of what the caller assumed + // when it started the future. + std::future pending_future; + + Poll poll_pending_future(); + void wait_pending_future(); +}; + +struct DeviceAccessControl: AccessControl { + // Number of currently-granted requests holding this exact location. Only + // locations with `alloc_count == 0` may sit in the LRU list / be evicted. + bool in_lru = false; + DeviceAccessControl* lru_prev = nullptr; + DeviceAccessControl* lru_next = nullptr; +}; + +struct DeviceLRU { + DeviceAccessControl* head = nullptr; + DeviceAccessControl* tail = nullptr; + + DeviceAccessControl* least_recently_used() const noexcept { + return head; + } + + void insert(DeviceAccessControl* loc) noexcept { + KMM_ASSERT(!loc->in_lru); + loc->lru_prev = tail; + loc->lru_next = nullptr; + + if (tail != nullptr) { + tail->lru_next = loc; + } else { + head = loc; + } + + tail = loc; + loc->in_lru = true; + } + + void remove(DeviceAccessControl* loc) noexcept { + if (!loc->in_lru) { + return; + } + + if (loc->lru_prev != nullptr) { + loc->lru_prev->lru_next = loc->lru_next; + } else { + head = loc->lru_next; + } + + if (loc->lru_next != nullptr) { + loc->lru_next->lru_prev = loc->lru_prev; + } else { + tail = loc->lru_prev; + } + + loc->lru_prev = nullptr; + loc->lru_next = nullptr; + loc->in_lru = false; + } +}; + +struct MemoryBufferImpl: reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(MemoryBufferImpl) + public: + const size_t size_in_bytes; + + MemoryBufferImpl( + std::string name, + bool evictable, + std::unique_ptr data, + std::optional home_memory_id = {} + ) : + size_in_bytes(data->size_in_bytes()), + name(std::move(name)), + evictable(evictable), + data(std::move(data)), + home_memory_id(home_memory_id) {} + + bool is_compatible(MemoryId memory_id, AccessKind mode) noexcept; + bool try_register_request(BufferQueueNode* req) noexcept; + void unregister_request(BufferQueueNode* req) noexcept; + + AccessControl& location(MemoryId id) noexcept { + if (id.is_host()) { + return host_location; + } else { + return device_locations[id.as_device().get()]; + } + } + + const AccessControl& location(MemoryId id) const noexcept { + if (id.is_host()) { + return host_location; + } else { + return device_locations[id.as_device().get()]; + } + } + + // A host location that's still `is_valid` while its producer is in flight counts as + // "found" here too: this only picks a copy *source*, and the eventual consumer + // (`poll_copy`/`do_copy`) waits out `pending_future` before actually reading it. Treating + // an in-flight producer as absent instead would make callers re-produce the data from + // scratch (redundant, and wrong for a non-idempotent producer) instead of reusing what's + // already in flight. + bool find_valid_location(MemoryId exclude, MemoryId& out) const noexcept { + // Prefer the buffer's home location, if it has one and it's still valid. + if (home_memory_id.has_value()) { + if (*home_memory_id != exclude && location(*home_memory_id).is_valid) { + out = *home_memory_id; + return true; + } + } + + if (host_location.is_valid && MemoryId::host() != exclude) { + out = MemoryId::host(); + return true; + } + + for (size_t id = 0; id < MAX_DEVICES; id++) { + if (device_locations[id].is_valid && MemoryId::device(DeviceId(id)) != exclude) { + out = MemoryId::device(DeviceId(id)); + return true; + } + } + + return false; + } + + MemoryId find_preferred_location(MemoryId fallback) { + if (is_valid(fallback)) { + return fallback; + } + + if (find_valid_location(fallback, fallback)) { + return fallback; + } + + if (home_memory_id.has_value()) { + return *home_memory_id; + } + + return fallback; + } + + // Strict: `true` if the data has a producer (either data is available or will be available). + bool is_valid(MemoryId id) { + return location(id).is_valid; + } + + bool is_allocated(MemoryId id) { + return location(id).is_allocated; + } + + AllocResult try_allocate_location( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId dst_id + ); + + bool allocate_host(MemorySystem& system, const DeviceStreamId& stream_hint); + bool deallocate_host(MemorySystem& system, const DeviceStreamId& stream_hint); + void increment_host_users() noexcept; + void decrement_host_users() noexcept; + + AllocResult try_allocate_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId id + ); + bool deallocate_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId id, + DeviceLRU& lru + ); + void increment_device_users(DeviceId id, DeviceLRU& lru) noexcept; + void decrement_device_users(DeviceId id, DeviceLRU& lru) noexcept; + + void evict_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId id, + DeviceLRU& lru + ); + Poll ensure_alloc_valid( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId memory_id + ); + void invalidate_other_allocs(MemoryId memory_id); + DeviceEventSet invalidate_all(); + + // Called once a request's location has been granted, right before the + // caller is allowed to actually read/write through it. + Poll before_access( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId memory_id, + AccessKind mode + ); + + // Called right after the caller is done reading/writing through a + // granted location, recording the resulting dependencies. + void after_access(MemoryId memory_id, AccessKind mode, const DeviceEventSet& deps); + + // Returns `Pending` (without touching any state) if `src_id` is host and its data is + // still being produced by an in-flight `pending_future`. + Poll poll_copy( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId src_id, + MemoryId dst_id + ); + void do_copy( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId src_id, + MemoryId dst_id + ); + + // Returns the accessor granting access to this buffer (once `before_access` has returned + // `Ready`), and inserts into `deps_out` the events that must complete before it is safe to + // read/write through it. Reads live off the location's current epoch/write/read events, so + // this may safely be called after `before_access` rather than only from within it. + BufferAccessor access(MemoryId memory_id, AccessKind mode, DeviceEventSet& deps_out); + + const std::string name; + + // If false, this buffer's device locations are never inserted into the LRU. + const bool evictable; + + std::unique_ptr data; + + HostAccessControl host_location; + DeviceAccessControl device_locations[MAX_DEVICES]; + + BufferQueueNode* queue_head = nullptr; + BufferQueueNode* queue_tail = nullptr; + + // The location of the first access ever granted to this buffer. Set once, by + // `before_access`, and never changed afterwards. + std::optional home_memory_id; + + // Set by `MemoryManager::release_buffer` before it tears down allocations, + // to guard against the buffer being used in a new transaction while that + // teardown is in progress. + bool released = false; +}; + +} // namespace kmm \ No newline at end of file diff --git a/include/kmm/runtime/memory_manager.hpp b/include/kmm/runtime/memory_manager.hpp index f0ef1600..d3cf56b4 100644 --- a/include/kmm/runtime/memory_manager.hpp +++ b/include/kmm/runtime/memory_manager.hpp @@ -1,94 +1,94 @@ #pragma once +#include +#include #include -#include -#include +#include -#include "kmm/core/buffer.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/data_interfaces/base.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/identifiers.hpp" #include "kmm/runtime/memory_system.hpp" -#include "kmm/runtime/stream_manager.hpp" +#include "kmm/utils/notify.hpp" #include "kmm/utils/poll.hpp" +#include "kmm/utils/refcnt_ptr.hpp" +#include "kmm/utils/small_vector.hpp" namespace kmm { -using TransactionId = uint64_t; +// MemoryRequestImpl / MemoryRequest are forward-declared in buffer.hpp (included above), since +// BufferRequest needs them too and this header would otherwise have to be included before it. +struct MemoryTransactionImpl; +struct MemoryBufferImpl; +struct MemoryRequestImpl; -class MemoryManager { - KMM_NOT_COPYABLE_OR_MOVABLE(MemoryManager) +KMM_REFCNT_TRAITS_FWD(MemoryTransactionImpl) +KMM_REFCNT_TRAITS_FWD(MemoryBufferImpl) +KMM_REFCNT_TRAITS_FWD(MemoryRequestImpl) + +using MemoryBuffer = refcnt_ptr; +using MemoryTransaction = refcnt_ptr; +using MemoryRequest = refcnt_ptr; + +enum struct AccessKind { ReadOnly, SharedWrite, Exclusive }; +class MemoryManager { public: - struct Request; - struct Buffer; - struct Device; - struct Transaction; + struct Impl; - MemoryManager(std::shared_ptr memory); + explicit MemoryManager(refcnt_ptr memory_system); ~MemoryManager(); - bool is_idle(DeviceStreamManager& streams) const; + MemoryBuffer create_buffer( + std::unique_ptr data, + std::string name, + bool evictable = true, + std::optional home_memory_id = {} + ); - std::shared_ptr create_transaction(std::shared_ptr parent = nullptr); + void release_buffer(MemoryBuffer buffer); - std::shared_ptr create_buffer(BufferLayout layout, std::string name = ""); - void delete_buffer(std::shared_ptr buffer); + /// Creates a new transaction, optionally as a child of `parent`. + MemoryTransaction create_transaction(MemoryTransaction parent = {}); - std::shared_ptr create_request( - std::shared_ptr buffer, + MemoryRequest create_request( + const MemoryBuffer& buffer, MemoryId memory_id, - AccessMode mode, - std::shared_ptr parent + AccessKind mode, + MemoryTransaction parent = {}, + NotifyHandle callback = {} ); - Poll poll_request(Request& req, DeviceEventSet& deps_out); - void release_request(std::shared_ptr req, DeviceEvent event = {}); - BufferAccessor get_accessor(Request& req); + Poll poll_request(const DeviceStreamId& stream_hint, const MemoryRequest& request); - private: - void allocate_host(Buffer& buffer, DeviceId device_affinity); - void deallocate_host(Buffer& buffer); - - bool try_free_device_memory(DeviceId device_id); - AllocationResult try_allocate_device_async(DeviceId device_id, Buffer& buffer); - void deallocate_device_async(DeviceId device_id, Buffer& buffer); + /// Returns the accessor for a request that has reached `Ready` (see `poll_request`). + BufferAccessor access_request(const MemoryRequest& request, DeviceEventSet& deps_out); - void lock_allocation_host(Buffer& buffer, DeviceId device_affinity, Request& req); - static void unlock_allocation_host(Buffer& buffer, Request& req); + void release_request(MemoryRequest request, const DeviceEventSet& deps = {}); - bool try_lock_allocation_device(DeviceId device_id, Buffer& buffer, Request& req); - void unlock_allocation_device(DeviceId device_id, Buffer& buffer, Request& req) noexcept; - - void prepare_access_to_buffer( + void prefetch_buffer( + const MemoryBuffer& buffer, MemoryId memory_id, - Buffer& buffer, - AccessMode mode, - DeviceEventSet& deps_out + AccessKind mode = AccessKind::ReadOnly ); - static void finalize_access_to_buffer( - MemoryId memory_id, - Buffer& buffer, - AccessMode mode, - DeviceEvent event - ) noexcept; - - static std::optional find_valid_device_entry(const Buffer& buffer); - void make_entry_valid(MemoryId memory_id, Buffer& buffer, DeviceEventSet& deps_out); - void make_entry_exclusive(MemoryId memory_id, Buffer& buffer, DeviceEventSet& deps_out); - - DeviceEvent copy_h2d(DeviceId device_id, Buffer& buffer); - DeviceEvent copy_d2h(DeviceId device_id, Buffer& buffer); - DeviceEvent copy_d2d(DeviceId device_src_id, DeviceId device_dst_id, Buffer& buffer); - - Device& device_at(DeviceId id) noexcept; - bool is_out_of_memory(DeviceId device_id, Request& req); - - void check_consistency() const; - - std::shared_ptr m_memory; - std::unique_ptr m_devices; - std::unordered_set> m_buffers; - std::unordered_set> m_active_requests; - uint64_t m_next_transaction_id = 1; - uint64_t m_next_request_id = 1; + + void try_evict_buffer(const MemoryBuffer& buffer, MemoryId memory_id); + + void invalidate_buffer(const MemoryBuffer& buffer); + + void trim_device(DeviceId id, size_t bytes_remaining = 0, bool evict = false); + + void make_progress(); + + private: + std::unique_ptr m_impl; }; +std::ostream& operator<<(std::ostream& stream, AccessKind access); + } // namespace kmm + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/runtime/memory_system.hpp b/include/kmm/runtime/memory_system.hpp index 3f4c2908..0b916a29 100644 --- a/include/kmm/runtime/memory_system.hpp +++ b/include/kmm/runtime/memory_system.hpp @@ -1,140 +1,211 @@ #pragma once -#include "kmm/runtime/allocators/base.hpp" -#include "kmm/runtime/stream_manager.hpp" -#include "kmm/utils/macros.hpp" +#include +#include +#include +#include +#include + +#include "runtime_config.hpp" + +#include "kmm/core/macros.hpp" +#include "kmm/runtime/allocators/device.hpp" +#include "kmm/runtime/allocators/pinned.hpp" +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/device_stream.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/system_info.hpp" +#include "kmm/utils/refcnt_ptr.hpp" namespace kmm { -class MemorySystem { +struct MemoryStats { + size_t bytes_inuse = 0; + size_t bytes_allocated = 0; + size_t max_bytes_inuse = 0; + size_t bytes_to_host = 0; + size_t bytes_to_device[MAX_DEVICES] {}; + + void record_allocation(size_t nbytes) { + bytes_allocated += nbytes; + bytes_inuse += nbytes; + max_bytes_inuse = std::max(max_bytes_inuse, bytes_inuse); + } + + void record_deallocation(size_t nbytes) { + bytes_inuse -= nbytes; + } +}; + +class MemorySystem: public reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(MemorySystem) + public: - virtual ~MemorySystem() = default; + MemorySystem( + const SystemInfo& system_info, + DeviceEventRegistry events, + const RuntimeConfig& config + ); - virtual void make_progress() {} - virtual void trim_host(size_t bytes_remaining = 0) {} - virtual void trim_device(size_t bytes_remaining = 0) {} + ~MemorySystem(); - virtual AllocationResult allocate_host( - size_t nbytes, - DeviceId device_affinity, + void make_progress(); + void trim_host(size_t bytes_remaining = 0); + void trim_device(DeviceId id, size_t bytes_remaining = 0); + + /// Real bytes the memory's allocator holds reserved from the OS/driver. + size_t bytes_reserved(MemoryId id) const; + + AllocResult allocate_host( + BufferLayout layout, void** ptr_out, + const DeviceStreamId& stream_hint, DeviceEventSet& deps_out - ) = 0; + ); + + void deallocate_host( + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); - virtual void deallocate_host(void* ptr, size_t nbytes, DeviceEventSet deps) = 0; + void* translate_host_pointer(DeviceId device_id, void* host_ptr) const; + + AllocResult allocate_managed( + BufferLayout layout, + void** ptr_out, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out + ); - virtual AllocationResult allocate_device( + void deallocate_managed( + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); + + void prefetch_managed( + MemoryId memory_id, + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps + ); + + AllocResult allocate_device( DeviceId device_id, - size_t nbytes, + BufferLayout layout, g_device_ptr_t* ptr_out, + const DeviceStreamId& stream_hint, DeviceEventSet& deps_out - ) = 0; + ); - virtual void deallocate_device( + void deallocate_device( DeviceId device_id, g_device_ptr_t ptr, - size_t nbytes, - DeviceEventSet deps - ) = 0; + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); - virtual DeviceEvent copy_host_to_device( + DeviceEvent copy_host_to_device( DeviceId device_id, const void* src_addr, g_device_ptr_t dst_addr, size_t nbytes, - DeviceEventSet deps - ) = 0; + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); - virtual DeviceEvent copy_device_to_host( + DeviceEvent copy_device_to_host( DeviceId device_id, g_device_ptr_t src_addr, void* dst_addr, size_t nbytes, - DeviceEventSet deps - ) = 0; + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); - virtual DeviceEvent copy_device_to_device( + DeviceEvent copy_device_to_device( DeviceId src_device, DeviceId dst_device, g_device_ptr_t src_addr, g_device_ptr_t dst_addr, size_t nbytes, - DeviceEventSet deps - ) = 0; - - virtual bool is_copy_supported(MemoryId src, MemoryId dst) { - return true; - } -}; - -class MemorySystemImpl: public MemorySystem { - KMM_NOT_COPYABLE_OR_MOVABLE(MemorySystemImpl) - - public: - MemorySystemImpl( - std::shared_ptr stream_manager, - std::vector device_contexts, - std::unique_ptr host_mem, - std::vector> device_mem + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in ); - ~MemorySystemImpl(); - - void make_progress(); - void trim_host(size_t bytes_remaining = 0); - void trim_device(size_t bytes_remaining = 0); - - AllocationResult allocate_host( - size_t nbytes, - DeviceId device_affinity, - void** ptr_out, - DeviceEventSet& deps_out - ) final; - void deallocate_host(void* ptr, size_t nbytes, DeviceEventSet deps) final; - - AllocationResult allocate_device( + AllocResult allocate_host_and_copy_from_device( + BufferLayout layout, + void** dst_addr, DeviceId device_id, - size_t nbytes, - g_device_ptr_t* ptr_out, - DeviceEventSet& deps_out - ) final; - - void deallocate_device( - DeviceId device_id, - g_device_ptr_t ptr, - size_t nbytes, - DeviceEventSet deps - ) final; + g_device_ptr_t src_addr, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEvent& dep_out + ); - DeviceEvent copy_host_to_device( + AllocResult allocate_device_and_copy_from_host( DeviceId device_id, + BufferLayout layout, + g_device_ptr_t* dst_addr, const void* src_addr, - g_device_ptr_t dst_addr, - size_t nbytes, - DeviceEventSet deps - ) final; + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEvent& dep_out + ); - DeviceEvent copy_device_to_host( + DeviceEvent fill_device( DeviceId device_id, - g_device_ptr_t src_addr, - void* dst_addr, - size_t nbytes, - DeviceEventSet deps - ) final; + g_device_ptr_t addr, + const FillDescription& description, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); - DeviceEvent copy_device_to_device( - DeviceId src_device_id, - DeviceId dst_device_id, + /// Fill host memory at `addr` according to `description`. Runs on a background thread since + /// this is a CPU-bound operation; the returned future becomes ready once `deps_in` have + /// completed and the fill has finished. + std::future fill_host( + void* addr, + const FillDescription& description, + const DeviceEventSet& deps_in + ); + + DeviceEvent reduce_device( + DeviceId device_id, g_device_ptr_t src_addr, g_device_ptr_t dst_addr, - size_t nbytes, - DeviceEventSet deps - ) final; + g_device_ptr_t scratch_addr, + const ReductionDescription& description, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in + ); + + bool is_copy_supported(MemoryId src, MemoryId dst) const noexcept; private: - struct Device; - std::shared_ptr m_streams; - std::unique_ptr m_host; - std::unique_ptr m_devices[MAX_DEVICES]; + struct DeviceState; + + DeviceState& device_state(DeviceId id) const; + DeviceId affinity_for_stream(const DeviceStreamId& stream_hint); + bool same_context(DeviceId device_id, const DeviceStreamId& stream_hint); + + DeviceEventRegistry m_events; + DeviceDataStreams m_streams; + std::unique_ptr m_host_allocator; + MemoryStats m_host_stats; + std::unique_ptr m_managed_allocator; + MemoryStats m_managed_stats; + std::array, MAX_DEVICES> m_devices; + size_t m_num_devices = 0; + bool m_peer_access[MAX_DEVICES][MAX_DEVICES] {}; }; + } // namespace kmm diff --git a/include/kmm/runtime/reduction_manager.hpp b/include/kmm/runtime/reduction_manager.hpp new file mode 100644 index 00000000..a0a296b2 --- /dev/null +++ b/include/kmm/runtime/reduction_manager.hpp @@ -0,0 +1,64 @@ +#pragma once + +#include + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memops/types.hpp" +#include "kmm/runtime/memory_manager.hpp" +#include "kmm/utils/poll.hpp" +#include "kmm/utils/refcnt_ptr.hpp" + +namespace kmm { + +class MemorySystem; +class ReductionStateImpl; +using ReductionState = refcnt_ptr; + +KMM_REFCNT_TRAITS_FWD(ReductionStateImpl) + +class ReductionJob; + +class ReductionManager { + public: + ReductionManager(MemoryManager& memory_manager, refcnt_ptr memory_system); + ~ReductionManager(); + + /// Initialize a new reduction for the given buffer. + ReductionState initialize_reduction( + MemoryBuffer home_buffer, + ReductionOp op, + DataType dtype, + size_t count + ); + + /// Acquire a new partial buffer for the given reduction. + MemoryBuffer acquire_partial( + ReductionState& reduction, + MemoryId memory_id, + const DeviceStreamId& stream_hint + ); + + /// Check that a contribution combining with `op` over `dtype` elements is compatible with the + /// parameters this reduction was opened with by `begin_reduction`. Throws `std::runtime_error` + /// on a mismatch, since the partials are folded using the reduction's own op/dtype and a + /// mismatch would silently produce a wrong result. + void check_compatible(const ReductionState& reduction, ReductionOp op, DataType dtype) const; + + /// Transition the reduction from open (i.e., partial buffers can still be acquired) to + /// active (i.e., the reduction will be performed). + void submit_reduction(ReductionState& reduction, MemoryId memory_id); + + bool is_submitted(const ReductionState& reduction); + + Poll poll_reduction(ReductionState& reduction); + + void release_reduction(ReductionState& reduction); + + private: + MemoryManager& m_memory_manager; + refcnt_ptr m_memory_system; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/resource.hpp b/include/kmm/runtime/resource.hpp new file mode 100644 index 00000000..b0bfbf43 --- /dev/null +++ b/include/kmm/runtime/resource.hpp @@ -0,0 +1,70 @@ +#pragma once + +#include +#include +#include + +#include "kmm/runtime/buffer.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/identifiers.hpp" +#include "kmm/runtime/memory_manager.hpp" +#include "kmm/runtime/runtime.hpp" + +namespace kmm { + +struct BufferRequest { + MemoryId memory_id; + BufferId buffer_id; + AccessMode mode; +}; + +class ResourceRequest { + public: + size_t add(MemoryId memory_id, BufferId buffer_id, AccessMode mode) { + size_t index = m_requests.size(); + m_requests.push_back({memory_id, buffer_id, mode}); + return index; + } + + private: + friend class Runtime; + + std::vector m_requests; +}; + +class ResourceGrant { + KMM_NOT_COPYABLE_OR_MOVABLE(ResourceGrant) + + public: + BufferAccessor accessor(size_t index) const noexcept { + return m_entries.at(index).accessor; + } + + const DeviceEventSet& dependencies() const noexcept { + return m_deps; + } + + const MemoryTransaction& transaction() const noexcept { + return m_transaction; + } + + private: + friend class Runtime; + + struct Entry { + BufferId buffer_id; + MemoryRequest request; + BufferAccessor accessor {}; + }; + + ResourceGrant(std::vector entries, DeviceEventSet deps, MemoryTransaction transaction) : + m_entries(std::move(entries)), + m_deps(std::move(deps)), + m_transaction(std::move(transaction)) {} + + std::vector m_entries; + DeviceEventSet m_deps; + MemoryTransaction m_transaction; +}; + +} // namespace kmm diff --git a/include/kmm/runtime/runtime.hpp b/include/kmm/runtime/runtime.hpp index de713324..da43fcba 100644 --- a/include/kmm/runtime/runtime.hpp +++ b/include/kmm/runtime/runtime.hpp @@ -1,75 +1,300 @@ #pragma once -#include +#include +#include +#include +#include -#include "kmm/core/config.hpp" -#include "kmm/core/system_info.hpp" +#include "runtime_config.hpp" + +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/memops/types.hpp" +#include "kmm/runtime/memory_manager.hpp" #include "kmm/runtime/memory_system.hpp" -#include "kmm/runtime/scheduler.hpp" -#include "kmm/runtime/task_graph.hpp" +#include "kmm/runtime/system_info.hpp" +#include "kmm/utils/function_ref.hpp" namespace kmm { -class Runtime: public std::enable_shared_from_this { - KMM_NOT_COPYABLE_OR_MOVABLE(Runtime) +class RuntimeImpl; +class ResourceRequest; +class ResourceGrant; +/// Top-level handle to the KMM runtime: owns the system's `SystemInfo`, `DeviceStreamRegistry`, +/// `MemorySystem`, and `MemoryManager`, and exposes the buffer/request operations built on top of +/// them. Cheap to copy (a `refcnt_ptr` to the shared `RuntimeImpl`). +class Runtime { public: - Runtime( - std::vector contexts, - std::shared_ptr stream_manager, - std::shared_ptr memory_system, - const RuntimeConfig& config + /** + * Returns the physical-memory backend used to allocate and move buffer data. + */ + MemorySystem& memory_system() noexcept; + + /** + * Returns the registry of device streams used to schedule and track work on the devices. + */ + DeviceEventRegistry& event_registry() noexcept; + + /** + * Returns information about the machine's topology (hosts, devices, memories). + */ + const SystemInfo& system_info() const noexcept; + + /** + * Poll the Runtime once. This will update bookkeeping that needs periodic updating and checks + * if certain events have finished. + * + * @return The time when the next poll should happen. Polling before this time is a noop. + */ + std::chrono::system_clock::time_point poll_once(); + + /** + * Repeatedly poll the runtime until the given callback return `true`. This is useful to wait + * for a certain event to complete while the runtime system can still make progress in the + * background. + * + * @param callback Called after each poll to check if we are done. + * @param deadline Determines the maximum timeout. + * @return Returns `true` if the callback return `true`, and `false` if the deadline expired. + */ + bool poll_until_completion( + function_ref callback, + std::chrono::system_clock::time_point deadline = + std::chrono::system_clock::time_point::max() ); - ~Runtime(); - - BufferId create_buffer(BufferLayout layout); - void delete_buffer(BufferId buffer_id, EventList deps = {}); - void check_buffer(BufferId id); - - bool query_event(EventId event_id, std::chrono::system_clock::time_point deadline); - bool is_idle(); - void trim_memory(); - void make_progress(); - void shutdown(); - - template> - R schedule(F fun) { - std::lock_guard guard {m_mutex}; - - if constexpr (std::is_void_v) { - auto stage = TaskGraph(&m_graph_state); - fun(stage); - this->commit_impl(stage); - } else { - auto stage = TaskGraph(&m_graph_state); - auto result = fun(stage); - this->commit_impl(stage); - return result; - } - } - - const SystemInfo& system_info() const { - return m_info; - } + + /** + * Register a caller-provided `DataInterface` as a buffer. This is the low-level primitive that + * `create_buffer` and `adopt_buffer` are built on. + * + * @param data The interface backing the buffer. Must not be null. + * @param name The name of the new buffer. + * @param home If set, the memory where the buffer is preferentially kept resident. + * @param evictable If false, the buffer's device locations are never evicted. + * @return The identifier of the new buffer. + */ + BufferId register_buffer( + std::unique_ptr data, + std::string name, + std::optional home = {}, + bool evictable = true + ); + + /** + * Create a new buffer in the runtime system. + * + * @param layout The layout of the new buffer. + * @param name The name of the new buffer. + * @param fill_value If non-empty, the buffer's contents are set to repeated copies of this + * value the first time it is materialized in any memory. + * @param home If set, the memory where the buffer is preferentially kept resident, used to + * pick a copy source once the buffer becomes valid there instead of the location of first + * access. + * @param kind The kind of storage backing the buffer (discrete, managed, or host-pinned). If + * empty, `RuntimeConfig::default_buffer_kind` is used. + * @return The identifier of the new buffer. + */ + BufferId create_buffer( + BufferLayout layout, + std::string name, + FillValue fill_value = {}, + std::optional home = {}, + std::optional kind = std::nullopt + ); + + /** + * Register a pre-existing, externally-owned allocation with the runtime as a buffer. KMM never + * allocates, frees, or relocates the memory: the buffer is pinned to `memory_id` and any + * attempt to access or copy it on another memory throws. This lets an application hand KMM a + * pointer it already owns (host or device) without a copy; the caller keeps responsibility for + * keeping the allocation alive for as long as the buffer is in use. + * + * @param layout The layout (size/alignment) describing `external_ptr`. + * @param name The name of the new buffer. + * @param external_ptr The pre-existing allocation. Must not be null and must be valid in + * `memory_id`. + * @param memory_id The memory that `external_ptr` lives in. + * @return The identifier of the new buffer. + */ + BufferId adopt_buffer( + BufferLayout layout, + std::string name, + void* external_ptr, + MemoryId memory_id + ); + + /** + * Release a buffer, freeing its memory once any pending accesses to it have completed. + * + * @param id The identifier of the buffer to release. + */ + void release_buffer(BufferId id); + + /** + * Puts a buffer into reduction mode: until `finalize_reduction` is called, the buffer may + * only be accessed in reduce mode (`Requisition::add_reduction`). Any other access is a + * bug and throws instead of being served. + * + * Throws if the buffer is already in reduction mode. + */ + void begin_reduction(BufferId id, DataType dtype, ReductionOp op); + + /** + * Finalizes the reduction previously started with `begin_reduction`, folding every value + * written to the buffer during reduce-mode access into it and returning the buffer to + * regular read/write access. The fold may still be pending when this call returns; `submit` + * waits for it automatically before granting the next read/write access. + * + * Throws if the buffer is not currently in reduction mode. + */ + void finalize_reduction(BufferId id, MemoryId memory_id = MemoryId::host()); + + /** + * Abandons the reduction previously started with `begin_reduction`: discards every value + * written to the buffer during reduce-mode access (without folding any of it in) and returns + * the buffer to regular read/write access. Unlike `finalize_reduction`, there is no fold to + * wait for, so this cannot fail asynchronously. A no-op if the buffer is not currently in + * reduction mode. + */ + void rollback_reduction(BufferId id); + + /** + * Ensure a valid copy of the buffer is available in the given memory, fetching it if needed. + * This is just a hint. The runtime system might ignore the request (for example, if out of + * memory or if the buffer is currently locked by another thread). + */ + void prefetch_buffer(BufferId id, MemoryId memory_id, bool invalidate_others = false); + + /** + * Mark a buffer as poisoned, so future accesses to it rethrow `reason` instead of succeeding. + * This is useful for cases where a write failed and the content of the buffer is now in an + * invalid state, making it impossible for others to read. + */ + void poison_buffer(BufferId id, std::exception_ptr reason) noexcept; + + /** + * Returns a memory that currently holds a valid copy of the buffer, if any. + */ + std::optional find_valid_memory(BufferId) const; + + /** + * Returns the buffer's home memory, if it has one. The home is either the memory passed to + * `create_buffer`, or (if none was passed) the memory of the first access ever granted to + * the buffer; it is `nullopt` only for a buffer that has never been created with an explicit + * home and has not yet been accessed. + */ + std::optional buffer_home(BufferId) const; + + /** + * Returns whether the given memory currently holds a valid copy of the buffer. + */ + bool is_valid(BufferId id, MemoryId memory_id); + + /** + * Returns whether the buffer currently has memory allocated for it in the given memory. + */ + bool is_allocated(BufferId id, MemoryId memory_id); + + /** + * Try to evict the buffer's copy from the given memory to free up space, if possible. + */ + void try_evict_buffer(BufferId id, MemoryId memory_id); + + /** + * Invalidate all copies of the buffer, discarding its contents. + */ + void invalidate_buffer(BufferId id); + + /** + * Release memory that the given memory's allocator is holding cached but not using, handing it + * back to the OS. Memory that is currently in use is not freed. If the allocator has no free + * memory available to give back or if it does not cache allocations, this is a no-op. + * + * @param memory_id The memory whose allocator pool should be trimmed. + * @param bytes_to_keep The amount of cached memory the allocator may keep reserved. Anything + * above this that is not in use is released. + * @param evict By default, the system only releases unused allocations from, for example, a + * memory pool. If this is `true`, then buffers are also forcefully evicted to make enough + * space to reach `bytes_to_keep`. + */ + void trim(MemoryId memory_id, size_t bytes_to_keep = 0, bool evict = false); + + /** + * Submit a batch of buffer requests. If a stream is provided, all required dependencies will + * be put onto the stream and this method returns immediately. If no stream is provided, the + * method blocks until the dependencies are available. Every request is granted by the time + * this returns; hand the result to `release` once the access is done. + */ + ResourceGrant submit( + ResourceRequest requests, + std::optional stream = std::nullopt, + MemoryTransaction parent = {} + ); + + /** + * Release a grant obtained from `submit`, allowing others to access its buffers once `deps` + * complete. `grant` is neither movable nor copyable, so it is taken by reference. + */ + void release(ResourceGrant& grant, DeviceEventSet deps = {}); + + /** + * Poisons every buffer `grant` accessed for write/reduce, so future accesses to them rethrow + * `reason` instead of succeeding. See `poison_buffer`. + */ + void poison(const ResourceGrant& grant, std::exception_ptr reason) noexcept; + + /** + * Block the calling thread until the given device event has completed. + */ + void synchronize(const DeviceEvent& e); + + /** + * Block the calling thread until all device events in the set have completed. + */ + void synchronize(const DeviceEventSet& e); + + /** + * Block the calling thread until all events on all streams have completed. + */ + void synchronize(); + + /** + * Copy from the given src buffer to the given dst buffer according to the given description. + * If a stream is provided, it is used as a hint for where to schedule the required transfers. + */ + DeviceEvent submit_copy( + BufferId dst_id, + BufferId src_id, + CopyDescription description, + MemoryId memory_id, + std::optional stream = std::nullopt, + MemoryTransaction parent = {} + ); + + /** + * Reduce from the given src buffer into the given dst buffer according to the given + * description. Both buffers are accessed in `memory_id`, fetching src there first if needed. + * If a stream is provided, it is used as a hint for where to schedule the required work. + */ + DeviceEvent submit_reduction( + BufferId dst_id, + BufferId src_id, + ReductionDescription description, + MemoryId memory_id, + std::optional stream = std::nullopt, + MemoryTransaction parent = {} + ); + + explicit Runtime(refcnt_ptr); private: - EventId commit_impl(TaskGraph& g); - void make_progress_impl(); - bool is_idle_impl(); - - mutable std::mutex m_mutex; - std::chrono::system_clock::time_point m_next_updated_planned = std::chrono::system_clock::now(); - mutable bool m_has_shutdown = false; - std::shared_ptr m_memory_system; - std::shared_ptr m_memory_manager; - std::shared_ptr m_buffer_registry; - std::shared_ptr m_stream_manager; - std::shared_ptr m_devices; - SystemInfo m_info; - Scheduler m_scheduler; - TaskGraphState m_graph_state; + refcnt_ptr m_impl; }; -std::shared_ptr make_worker(const RuntimeConfig& config); +KMM_REFCNT_TRAITS_FWD(RuntimeImpl) + +Runtime make_runtime(const RuntimeConfig& config = default_config_from_environment()); -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/core/config.hpp b/include/kmm/runtime/runtime_config.hpp similarity index 71% rename from include/kmm/core/config.hpp rename to include/kmm/runtime/runtime_config.hpp index af9baae1..5a3c5981 100644 --- a/include/kmm/core/config.hpp +++ b/include/kmm/runtime/runtime_config.hpp @@ -2,6 +2,7 @@ #include #include +#include namespace kmm { @@ -27,6 +28,20 @@ enum struct DeviceMemoryKind { NoPool, }; +enum struct BufferKind { + /// One separate allocation per memory, with explicit copies between them when the buffer is + /// accessed on a memory that does not hold a valid copy. This is the default. + Discrete, + + /// A single CUDA/HIP managed-memory allocation (`cudaMallocManaged`). The driver migrates + /// pages between host and devices on demand; KMM never issues explicit copies. + Managed, + + /// A single pinned host allocation. Devices access it directly (zero-copy) through a mapped + /// pointer; KMM never stages a device-side copy. + HostPinned, +}; + struct RuntimeConfig { /// The type of memory pool to use for the host. HostMemoryKind host_memory_kind = HostMemoryKind::NoPool; @@ -37,7 +52,7 @@ struct RuntimeConfig { /// If nonzero, use an arena allocator on the host. This will allocate large blocks of the /// specified size, which are further split into smaller allocations by the KMM runtime system. /// This reduces the number of memory allocation requests to the OS. - size_t host_memory_block_size = 0; + size_t host_memory_block_size = size_t(1024) * 1024 * 100; /// The type of memory pool to use on the GPU. DeviceMemoryKind device_memory_kind = DeviceMemoryKind::DefaultPool; @@ -54,14 +69,17 @@ struct RuntimeConfig { /// If nonzero, use an arena allocator on each device. This will allocate large blocks of the /// specified size, from which smaller allocations are subsequently sub-allocated. - size_t device_memory_block_size = 0; + size_t device_memory_block_size = size_t(1024) * 1024 * 500; - /// The number of concurrent streams on each device for execution of kernels. - size_t device_concurrent_streams = 8; + /// The number of concurrent streams on each device for data transfers. + size_t device_concurrent_streams = 4; /// Enable this run the system in debug mode. This will be significantly slower, but can be /// used to track down synchronization bugs. bool debug_mode = false; + + /// The buffer kind used by `Runtime::create_buffer` when it is called without an explicit kind. + BufferKind default_buffer_kind = BufferKind::Discrete; }; RuntimeConfig default_config_from_environment(); diff --git a/include/kmm/runtime/scheduler.hpp b/include/kmm/runtime/scheduler.hpp deleted file mode 100644 index f309da6f..00000000 --- a/include/kmm/runtime/scheduler.hpp +++ /dev/null @@ -1,104 +0,0 @@ -#pragma once - -#include - -#include "kmm/core/commands.hpp" -#include "kmm/core/resource.hpp" -#include "kmm/runtime/buffer_registry.hpp" -#include "kmm/runtime/device_resources.hpp" -#include "kmm/runtime/memory_manager.hpp" -#include "kmm/runtime/stream_manager.hpp" -#include "kmm/runtime/task.hpp" -#include "kmm/utils/poll.hpp" - -namespace kmm { - -struct SchedulerQueue; - -class TaskRecord { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskRecord) - - friend class Scheduler; - enum struct Status { // - Init, - AwaitingDependencies, - ReadyToStart, - Running, - WaitingForCompletion, - Completed - }; - - public: - TaskRecord(EventId event_id, std::unique_ptr task) : - event_id(event_id), - task(std::move(task)) {} - - EventId id() const { - return event_id; - } - - private: - EventId event_id; - Status status = Status::Init; - SchedulerQueue* queue = nullptr; - - EventList predecessors; - size_t predecessors_pending = 0; - DeviceEventSet input_events; - - small_vector, 4> successors; - DeviceEventSet output_events; - - std::shared_ptr next = nullptr; - std::unique_ptr task = nullptr; -}; - -class Scheduler { - KMM_NOT_COPYABLE_OR_MOVABLE(Scheduler) - - public: - Scheduler( - std::shared_ptr device_resources, - std::shared_ptr stream_manager, - std::shared_ptr buffer_registry, - bool debug_mode - ); - - ~Scheduler(); - - void submit(EventId event_id, std::unique_ptr task, EventList dependencies); - bool is_completed(EventId event_id) const; - bool is_idle() const; - void make_progress(); - - DeviceResources& devices() { - return *m_device_resources; - } - - DeviceStreamManager& streams() { - return *m_stream_manager; - } - - BufferRegistry& buffers() { - return *m_buffer_registry; - } - - private: - static size_t determine_queue_id(const Task&); - void enqueue_if_ready(const TaskRecord* predecessor, const std::shared_ptr& task); - std::shared_ptr dequeue_ready_task(); - void start_task(std::shared_ptr record); - Poll poll_completion(TaskRecord& record); - - std::shared_ptr m_device_resources; - std::shared_ptr m_stream_manager; - std::shared_ptr m_buffer_registry; - - std::vector m_ready_queues; - std::unordered_map> m_tasks; - std::shared_ptr m_running_head = nullptr; - TaskRecord* m_running_tail = nullptr; - bool m_debug_mode = false; -}; - -} // namespace kmm diff --git a/include/kmm/runtime/stream_manager.hpp b/include/kmm/runtime/stream_manager.hpp deleted file mode 100644 index 4b07a5d3..00000000 --- a/include/kmm/runtime/stream_manager.hpp +++ /dev/null @@ -1,247 +0,0 @@ -#pragma once - -#include -#include - -#include "kmm/core/identifiers.hpp" -#include "kmm/utils/gpu_utils.hpp" -#include "kmm/utils/notify.hpp" -#include "kmm/utils/small_vector.hpp" - -namespace kmm { - -class DeviceStreamManager; -class DeviceStream; -class DeviceEvent; -class DeviceEventSet; - -class DeviceStreamManager { - KMM_NOT_COPYABLE_OR_MOVABLE(DeviceStreamManager) - - public: - DeviceStreamManager(); - ~DeviceStreamManager(); - - bool make_progress(); - bool make_progress_for_stream(DeviceStream stream_index); - - DeviceStream create_stream(GPUContextHandle context, bool high_priority = false); - DeviceStream get_or_add_stream(GPUContextHandle context, g_stream_t stream); - - void wait_until_idle() const; - void wait_until_ready(DeviceStream stream) const; - void wait_until_ready(DeviceEvent event) const; - void wait_until_ready(const DeviceEventSet& events) const; - - bool is_idle() const; - bool is_ready(DeviceStream stream) const noexcept; - bool is_ready(DeviceEvent event) const noexcept; - bool is_ready(const DeviceEventSet& events) const noexcept; - bool is_ready(DeviceEventSet& events) const noexcept; - - void attach_callback(DeviceEvent event, NotifyHandle callback); - void attach_callback(DeviceStream stream, NotifyHandle callback); - - DeviceEvent record_event(DeviceStream stream); - void wait_on_default_stream(DeviceStream stream); - - void wait_for_event(DeviceStream stream, DeviceEvent event) const; - void wait_for_events(DeviceStream stream, const DeviceEventSet& events) const; - void wait_for_events(DeviceStream stream, const DeviceEvent* begin, const DeviceEvent* end) - const; - void wait_for_events(DeviceStream stream, const std::vector& events) const; - - /** - * Check if the given `source` event must occur before the given `target` event. In other words, - * if this function returns true, then `source` must be triggered before `target` can trigger. - */ - static bool event_happens_before(DeviceEvent source, DeviceEvent target); - - GPUContextHandle context(DeviceStream stream) const; - g_stream_t get(DeviceStream stream) const; - - template - DeviceEvent with_stream(DeviceStream stream, const DeviceEventSet& deps, F fun); - - template - DeviceEvent with_stream(DeviceStream stream, F fun); - - struct StreamState; - struct EventPool; - - private: - std::vector m_streams; - std::vector m_event_pools; -}; - -class DeviceStream { - public: - using index_type = uint8_t; - - DeviceStream(index_type i = 0) : m_index(i) {} - - index_type get() const { - return m_index; - } - - operator index_type() const { - return m_index; - } - - friend std::ostream& operator<<(std::ostream&, const DeviceStream& e); - - private: - index_type m_index; -}; - -class DeviceEvent { - public: - using index_type = uint64_t; - static constexpr index_type max_index = index_type(1) << 56; - - DeviceEvent() = default; - - DeviceEvent(DeviceStream stream, index_type index) noexcept { - KMM_ASSERT(index < max_index); - m_event_and_stream = (index_type(stream.get()) * max_index) + index; - } - - bool is_null() const noexcept { - return m_event_and_stream == 0; - } - - DeviceStream stream() const noexcept { - return static_cast(m_event_and_stream / max_index); - } - - index_type index() const noexcept { - return m_event_and_stream % max_index; - } - - constexpr bool operator==(const DeviceEvent& that) const noexcept { - return this->m_event_and_stream == that.m_event_and_stream; - } - - constexpr bool operator<(const DeviceEvent& that) const noexcept { - // This is equivalent to tuple(this.stream, this.event) < tuple(that.stream, that.event) - return this->m_event_and_stream < that.m_event_and_stream; - } - - KMM_IMPL_COMPARISON_OPS(DeviceEvent) - - friend std::ostream& operator<<(std::ostream&, const DeviceEvent& e); - - private: - uint64_t m_event_and_stream = 0; -}; - -class DeviceEventSet { - public: - DeviceEventSet() = default; - DeviceEventSet(const DeviceEventSet&) = default; - DeviceEventSet(DeviceEventSet&&) noexcept = default; - DeviceEventSet(std::initializer_list); - DeviceEventSet(DeviceEvent); - - DeviceEventSet& operator=(const DeviceEventSet&) = default; - DeviceEventSet& operator=(DeviceEventSet&&) noexcept = default; - DeviceEventSet& operator=(std::initializer_list); - - /** - * Insert the given event into the set. - */ - void insert(DeviceEvent e) noexcept; - - /** - * Insert all events from `that` into this set. - */ - void insert(const DeviceEventSet& that) noexcept; - - /** - * Insert all events from `that` into this set. - */ - void insert(DeviceEventSet&& that) noexcept; - - /** - * Remove all events from the set for which the manager indicates that they are ready. - * - * @return true if all events were ready, false otherwise. - */ - bool remove_ready(const DeviceStreamManager&) noexcept; - - /** - * Remove events from the set for which the manager indicates that they are ready. This - * differs from `remove_ready` in that it only remove sthe events at the end of the list and - * does not reorder the leading events in the list. - * - * @return true if all events were ready, false otherwise. - */ - bool remove_ready_trailing(const DeviceStreamManager&) noexcept; - - /** - * Find the events from this set that have the same context as indicated by `context. The - * matching events are removed from the current set and returned as a new set. - * - * @return The extracted events. - */ - DeviceEventSet extract_events_for_context( - const DeviceStreamManager& manager, - GPUContextHandle context - ); - - /** - * Remove all events. - */ - void clear() noexcept; - - /** - * Returns `true` if the set is empty, `false` otherwise. - */ - bool is_empty() const noexcept; - - /** - * Returns pointer to the first event. - */ - const DeviceEvent* begin() const noexcept; - - /** - * Returns pointer to one past the last event. - */ - const DeviceEvent* end() const noexcept; - - friend DeviceEventSet operator|(const DeviceEventSet& a, const DeviceEventSet& b) noexcept; - friend std::ostream& operator<<(std::ostream&, const DeviceEventSet& e); - - private: - small_vector m_events; -}; - -template -DeviceEvent DeviceStreamManager::with_stream( - DeviceStream stream, - const DeviceEventSet& deps, - F fun -) { - wait_for_events(stream, deps); - return with_stream(stream, std::move(fun)); -} - -template -DeviceEvent DeviceStreamManager::with_stream(DeviceStream stream, F fun) { - try { - fun(get(stream)); - return record_event(stream); - } catch (...) { - wait_until_ready(stream); - throw; - } -} - -} // namespace kmm - -template<> -struct fmt::formatter: fmt::ostream_formatter {}; -template<> -struct fmt::formatter: fmt::ostream_formatter {}; -template<> -struct fmt::formatter: fmt::ostream_formatter {}; \ No newline at end of file diff --git a/include/kmm/core/system_info.hpp b/include/kmm/runtime/system_info.hpp similarity index 50% rename from include/kmm/core/system_info.hpp rename to include/kmm/runtime/system_info.hpp index 6585ed04..3ceaba83 100644 --- a/include/kmm/core/system_info.hpp +++ b/include/kmm/runtime/system_info.hpp @@ -1,22 +1,26 @@ #pragma once -#include +#include #include +#include #include -#include "kmm/core/identifiers.hpp" +#include "kmm/core/macros.hpp" +#include "kmm/runtime/identifiers.hpp" #include "kmm/utils/gpu_utils.hpp" namespace kmm { +/// Static hardware properties of a single CUDA device, plus its (primary) CUDA context -- in kmm, +/// a "device" and its primary CUDA context are the same thing. class DeviceInfo { public: static constexpr size_t NUM_ATTRIBUTES = G_DEVICE_ATTRIBUTE_MAX; - DeviceInfo(DeviceId id, GPUContextHandle context, size_t num_concurrent_streams = 1); + DeviceInfo(DeviceId id, g_context_t context); /** - * Returns the name of the device as provided by `g_device_get_name`. + * Returns the name of the device as provided by `gpuDeviceGetName`. */ std::string name() const { return m_name; @@ -26,7 +30,7 @@ class DeviceInfo { * Returns which memory this device has affinity to. */ MemoryId memory_id() const { - return MemoryId(m_id); + return MemoryId::device(m_id); } /** @@ -37,10 +41,24 @@ class DeviceInfo { } /** - * Return this device as a `g_device_t`. + * Return this device as a `GPUdevice`. */ g_device_t device_ordinal() const { - return m_device_id; + return m_device; + } + + /** + * Return the (primary) CUDA context of this device. + */ + g_context_t context() const { + return m_context; + } + + /** + * Return the (primary) CUDA context of this device. + */ + GPUContextId context_id() const { + return m_context_id; } /** @@ -61,10 +79,9 @@ class DeviceInfo { dim3 max_grid_dim() const; /** - * Returns the compute capability of this device as integer `MAJOR * 10 + MINOR` (For example, - * `86` means capability 8.6) + * Returns the compute capability of this device as integer (Major, Minor) */ - int compute_capability() const; + std::pair compute_capability() const; /** * Returns the maximum number of threads per block supported by this device. @@ -76,24 +93,33 @@ class DeviceInfo { */ int attribute(g_device_attribute_t attrib) const; - size_t num_compute_streams() const { - return m_concurrent_stream; - } - private: DeviceId m_id; + g_device_t m_device; + g_context_t m_context; + GPUContextId m_context_id; std::string m_name; - g_device_t m_device_id; size_t m_memory_capacity; - size_t m_concurrent_stream = 1; - std::array m_attributes; + size_t m_total_memory; + int m_compute_capability_major; + int m_compute_capability_minor; }; +/// Static hardware topology of the machine: the number of CUDA devices and their properties. +/// Queried once (via the CUDA driver API) at construction and never changes afterwards. At most +/// `MAX_DEVICES` devices are reported, even if more are physically present, since `DeviceId` +/// itself cannot represent more. +/// +/// Retains each device's primary CUDA context for the lifetime of this object (released again on +/// destruction), so this also doubles as the map from a CUDA context back to its `DeviceId` -- +/// see `device_id`. class SystemInfo { + KMM_NOT_COPYABLE_OR_MOVABLE(SystemInfo) + public: - SystemInfo() = default; + SystemInfo(); SystemInfo(std::vector devices); - SystemInfo(SystemInfo info, std::vector subresources); + ~SystemInfo(); /** * Returns the number of GPUs in the system. @@ -108,22 +134,7 @@ class SystemInfo { /** * Find the device that has the given device ordinal. */ - const DeviceInfo& device_by_ordinal(g_device_t device_ordinal) const; - - /** - * Return a list of the available processors in the system. - */ - std::vector resources() const; - - /** - * Return a list of the available memories in the system. - */ - std::vector memories() const; - - /** - * Returns the highest affinity memory for the given processor. - */ - MemoryId affinity_memory(ResourceId proc_id) const; + const DeviceInfo& device_by_ordinal(g_device_t ordinal) const; /** * Returns the highest affinity memory for the given device. @@ -131,23 +142,17 @@ class SystemInfo { MemoryId affinity_memory(DeviceId device_id) const; /** - * Returns the processor that has the highest affinity for accessing the given memory. - */ - static ResourceId affinity_processor(MemoryId memory_id); - - /** - * Checks if the given processor can access the given memory. + * Returns the device belonging to the given context. */ - bool is_memory_accessible(MemoryId memory_id, ResourceId proc_id) const; + const DeviceInfo& device_from_context(g_context_t context) const; /** - * Checks if the given device can access the given memory. + * Returns the device belonging to the given stream. */ - bool is_memory_accessible(MemoryId memory_id, DeviceId device_id) const; + const DeviceInfo& device_from_stream(g_stream_t stream) const; private: std::vector m_devices; - std::vector m_resources; }; -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/include/kmm/runtime/task.hpp b/include/kmm/runtime/task.hpp deleted file mode 100644 index c7abc523..00000000 --- a/include/kmm/runtime/task.hpp +++ /dev/null @@ -1,259 +0,0 @@ -#pragma once - -#include - -#include "kmm/runtime/buffer_registry.hpp" -#include "kmm/runtime/device_resources.hpp" -#include "kmm/utils/poll.hpp" - -namespace kmm { - -class Scheduler; -class TaskRecord; - -class Task { - KMM_NOT_COPYABLE_OR_MOVABLE(Task) - - public: - Task() = default; - virtual ~Task() = default; - virtual void start(const DeviceEventSet& input_events) = 0; - virtual Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) = 0; - - virtual const char* name() const { - return typeid(*this).name(); - } -}; - -class JoinTask final: public Task { - public: - void start(const DeviceEventSet& input_events) final; - Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) final; - - private: - DeviceEventSet m_dependencies; -}; - -class DeleteBufferTask: public Task { - public: - DeleteBufferTask(BufferId buffer_id) : m_buffer_id(buffer_id) {} - - void start(const DeviceEventSet& input_events) final; - Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) final; - - private: - BufferId m_buffer_id; - DeviceEventSet m_dependencies; -}; - -class HostTask: public Task { - public: - HostTask(std::vector buffers) : m_buffers(std::move(buffers)) {} - - void start(const DeviceEventSet& input_events) final; - Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) final; - - protected: - virtual std::future submit( - Scheduler& scheduler, - std::vector accessors - ) = 0; - DeviceEventSet m_dependencies; - - private: - enum struct Status { - Init, - CreateBuffers, - PollingBuffers, - PollingDependencies, - Running, - Completing, - Completed - }; - - Status m_status = Status::Init; - std::future m_future; - std::vector m_buffers; - BufferRequestList m_requests; -}; - -class ExecuteHostTask: public HostTask { - public: - ExecuteHostTask( // - std::unique_ptr compute_task, - std::vector buffers - ) : - HostTask(std::move(buffers)), - m_task(std::move(compute_task)) {} - - std::future submit(Scheduler& scheduler, std::vector accessors) override; - - private: - std::unique_ptr m_task; -}; - -class CopyHostTask: public HostTask { - public: - CopyHostTask(BufferId src_buffer, BufferId dst_buffer, CopyDef definition) : - HostTask( - {BufferRequirement {src_buffer, MemoryId::host(), AccessMode::Read}, - BufferRequirement {dst_buffer, MemoryId::host(), AccessMode::ReadWrite}} - ), - m_copy(definition) {} - - std::future submit(Scheduler& scheduler, std::vector accessors) override; - - private: - CopyDef m_copy; -}; - -class ReductionHostTask: public HostTask { - public: - ReductionHostTask(BufferId src_buffer, BufferId dst_buffer, ReductionDef definition) : - HostTask( - {BufferRequirement {src_buffer, MemoryId::host(), AccessMode::Read}, - BufferRequirement {dst_buffer, MemoryId::host(), AccessMode::ReadWrite}} - ), - m_reduction(definition) {} - - std::future submit(Scheduler& scheduler, std::vector accessors) override; - - private: - ReductionDef m_reduction; -}; - -class FillHostTask: public HostTask { - public: - FillHostTask(BufferId dst_buffer, FillDef definition) : - HostTask({BufferRequirement {dst_buffer, MemoryId::host(), AccessMode::ReadWrite}}), - m_fill(definition) {} - - std::future submit(Scheduler& scheduler, std::vector accessors) override; - - private: - FillDef m_fill; -}; - -class DeviceTask: public Task, public DeviceResourceOperation { - public: - DeviceTask(ResourceId resource_id, std::vector buffers) : - m_resource(resource_id), - m_buffers(std::move(buffers)) {} - - void start(const DeviceEventSet& input_events) final; - Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) final; - - ResourceId resource_id() const { - return m_resource; - } - - private: - enum struct Status { - Init, - CreateBuffers, - PollingBuffers, - PollingDependencies, - Running, - Completing, - Completed - }; - - Status m_status = Status::Init; - ResourceId m_resource; - std::vector m_buffers; - BufferRequestList m_requests; - DeviceEvent m_execution_event; - DeviceEventSet m_dependencies; - DeviceEventSet m_local_dependencies; -}; - -class ExecuteDeviceTask: public DeviceTask { - public: - ExecuteDeviceTask( - ResourceId device_id, - std::unique_ptr compute_task, - std::vector buffers - ) : - DeviceTask(device_id, std::move(buffers)), - m_task(std::move(compute_task)) {} - - void execute(DeviceResource& device, std::vector accessors) final; - - private: - std::unique_ptr m_task; -}; - -class CopyDeviceTask: public DeviceTask { - public: - CopyDeviceTask( - DeviceId device_id, - BufferId src_buffer, - BufferId dst_buffer, - CopyDef definition - ) : - DeviceTask( - device_id, - {BufferRequirement {src_buffer, device_id, AccessMode::Read}, - BufferRequirement {dst_buffer, device_id, AccessMode::ReadWrite}} - ), - m_copy(definition) {} - - void execute(DeviceResource& device, std::vector accessors) final; - - private: - CopyDef m_copy; -}; - -class ReductionDeviceTask: public DeviceTask { - public: - ReductionDeviceTask( - DeviceId device_id, - BufferId src_buffer, - BufferId dst_buffer, - ReductionDef definition - ) : - DeviceTask( - device_id, - {BufferRequirement {src_buffer, device_id, AccessMode::Read}, - BufferRequirement {dst_buffer, device_id, AccessMode::ReadWrite}} - ), - m_reduction(std::move(definition)) {} - - void execute(DeviceResource& device, std::vector accessors) final; - - private: - ReductionDef m_reduction; -}; - -class FillDeviceTask: public DeviceTask { - public: - FillDeviceTask(DeviceId device_id, BufferId dst_buffer, FillDef definition) : - DeviceTask(device_id, {BufferRequirement {dst_buffer, device_id, AccessMode::ReadWrite}}), - m_fill(std::move(definition)) {} - - void execute(DeviceResource& device, std::vector accessors) final; - - private: - FillDef m_fill; -}; - -class PrefetchTask: public Task { - public: - PrefetchTask(BufferId buffer_id, MemoryId memory_id) : - m_buffers {{buffer_id, memory_id, AccessMode::Read}} {} - - void start(const DeviceEventSet& input_events) final; - Poll poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) final; - - private: - enum struct Status { Init, Polling, Completing, Completed }; - - Status m_status = Status::Init; - std::vector m_buffers; - BufferRequestList m_requests; - DeviceEventSet m_dependencies; -}; - -std::unique_ptr build_task_for_command(Command&& command); - -} // namespace kmm diff --git a/include/kmm/runtime/task_graph.hpp b/include/kmm/runtime/task_graph.hpp deleted file mode 100644 index b9987756..00000000 --- a/include/kmm/runtime/task_graph.hpp +++ /dev/null @@ -1,69 +0,0 @@ -#pragma once - -#include -#include -#include -#include - -#include "kmm/core/buffer.hpp" -#include "kmm/core/commands.hpp" -#include "kmm/core/reduction.hpp" -#include "kmm/utils/macros.hpp" - -namespace kmm { - -class TaskGraphState; - -class TaskGraph { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskGraph) - - public: - friend class TaskGraphState; - - struct Node { - EventId id; - Command command; - EventList dependencies; - }; - - TaskGraph(TaskGraphState* state); - - BufferId create_buffer(BufferLayout layout); - EventId delete_buffer(BufferId id, EventList deps = {}); - EventId insert_barrier(); - EventId insert_compute_task( - ResourceId process_id, - std::unique_ptr task, - std::vector buffers, - EventList deps = {} - ); - - EventId join_events(EventList deps); - EventId insert_node(Command command, EventList deps = {}); - - private: - TaskGraphState* m_state; - EventList m_events_since_last_barrier; - std::vector m_staged_nodes; - std::vector> m_staged_buffers; -}; - -class TaskGraphState { - KMM_NOT_COPYABLE_OR_MOVABLE(TaskGraphState) - public: - friend class TaskGraph; - - TaskGraphState() = default; - EventId commit( - TaskGraph& g, - std::vector& staged_nodes, - std::vector>& staged_buffers - ); - - private: - BufferId m_next_buffer_id = BufferId(1); - EventId m_next_event_id = EventId(1); - EventId m_last_barrier_id = EventId(1); -}; - -} // namespace kmm diff --git a/include/kmm/utils/backends.hpp b/include/kmm/utils/backends.hpp new file mode 100644 index 00000000..d66e79e9 --- /dev/null +++ b/include/kmm/utils/backends.hpp @@ -0,0 +1,9 @@ +#pragma once + +#ifdef KMM_USE_CUDA + #include "kmm/utils/backends/cuda.hpp" +#elif KMM_USE_HIP + #include "kmm/utils/backends/hip.hpp" +#else + #include "kmm/utils/backends/cpu.hpp" +#endif diff --git a/include/kmm/backends/cpu.hpp b/include/kmm/utils/backends/cpu.hpp similarity index 89% rename from include/kmm/backends/cpu.hpp rename to include/kmm/utils/backends/cpu.hpp index 7a3649d4..d839389d 100644 --- a/include/kmm/backends/cpu.hpp +++ b/include/kmm/utils/backends/cpu.hpp @@ -2,7 +2,7 @@ #include -#include "kmm/utils/macros.hpp" +#include "kmm/core/macros.hpp" namespace kmm { @@ -92,19 +92,24 @@ enum g_pointer_attribute_t {}; #define G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z g_device_attribute_t(7) #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR g_device_attribute_t(75) #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR g_device_attribute_t(76) +#define GPU_ERROR_PEER_ACCESS_ALREADY_ENABLED 704 g_result_t g_init(unsigned int); g_result_t g_device_get_count(int*); g_result_t g_device_get(g_device_t*, int); g_result_t g_device_get_name(char*, int, g_device_t); +g_result_t g_device_total_mem(size_t*, g_device_t); g_result_t g_device_get_attribute(int*, g_device_attribute_t, g_device_t); +g_result_t g_device_can_access_peer(int*, g_device_t, g_device_t); g_result_t g_ctx_get_device(g_device_t*); +g_result_t g_ctx_get_id(g_context_t, unsigned long long*); g_result_t g_ctx_create(g_context_t*, unsigned int, g_device_t); g_result_t g_ctx_destroy(g_context_t); g_result_t g_device_primary_ctx_retain(g_context_t*, g_device_t); g_result_t g_device_primary_ctx_release(g_device_t); g_result_t g_ctx_push_current(g_context_t); g_result_t g_ctx_pop_current(g_context_t*); +g_result_t g_ctx_enable_peer_access(g_context_t, unsigned int); // Stream Management Constants & Functions @@ -112,7 +117,11 @@ g_result_t g_ctx_pop_current(g_context_t*); #define G_EVENT_WAIT_DEFAULT g_event_wait_flags_t(0) g_result_t g_ctx_get_stream_priority_range(int*, int*); +g_result_t g_stream_create(g_stream_t*, unsigned int); g_result_t g_stream_create_with_priority(g_stream_t*, unsigned int, int); +g_result_t g_stream_get_device(g_stream_t, g_device_t*); +g_result_t g_stream_get_id(g_stream_t, unsigned long long*); +g_result_t g_stream_get_ctx(g_stream_t, g_context_t*); g_result_t g_stream_query(g_stream_t); g_result_t g_stream_synchronize(g_stream_t); g_result_t g_stream_destroy(g_stream_t); @@ -141,6 +150,7 @@ gpu_error_t gpu_event_elapsed_time(float* ms, g_event_t start, g_event_t stop); #define G_MEMORYTYPE_DEVICE G_MEMORYTYPE_DEVICE #define G_POINTER_ATTRIBUTE_MEMORY_TYPE g_pointer_attribute_t(2) #define G_POINTER_ATTRIBUTE_DEVICE_ORDINAL g_pointer_attribute_t(9) +#define G_MEM_ATTACH_GLOBAL 1 g_result_t g_mem_get_info(size_t*, size_t*); g_result_t gpu_mem_get_info(size_t*, size_t*); @@ -151,6 +161,9 @@ gpu_error_t gpu_free(g_device_ptr_t); g_result_t g_mem_host_alloc(void**, size_t, unsigned int); g_result_t g_mem_free_host(void*); g_result_t g_pointer_get_attribute(void*, g_pointer_attribute_t, g_device_ptr_t); +g_result_t g_mem_alloc_managed(g_device_ptr_t*, size_t, unsigned int); +g_result_t g_mem_host_get_device_pointer(g_device_ptr_t*, void*, unsigned int); +g_result_t g_mem_prefetch_async(g_device_ptr_t, size_t, int, g_stream_t); // Memory Copy Operations @@ -191,6 +204,10 @@ gpu_error_t gpu_memset_async(void* dst, int value, size_t sizeBytes, g_stream_t g_result_t g_memcpy_2d(const gpu_memcpy2d_t*); g_result_t g_memcpy_2d_async(const gpu_memcpy2d_t*, g_stream_t); +g_result_t g_memset_d2d8_async(g_device_ptr_t, size_t, unsigned char, size_t, size_t, g_stream_t); +g_result_t g_memset_d2d16_async(g_device_ptr_t, size_t, unsigned short, size_t, size_t, g_stream_t); +g_result_t g_memset_d2d32_async(g_device_ptr_t, size_t, unsigned int, size_t, size_t, g_stream_t); + // Memory Pool Management Constants & Functions #define G_MEM_ALLOCATION_TYPE_PINNED G_MEM_ALLOCATION_TYPE_PINNED diff --git a/include/kmm/backends/cuda.hpp b/include/kmm/utils/backends/cuda.hpp similarity index 80% rename from include/kmm/backends/cuda.hpp rename to include/kmm/utils/backends/cuda.hpp index 29c55936..6d6b793e 100644 --- a/include/kmm/backends/cuda.hpp +++ b/include/kmm/utils/backends/cuda.hpp @@ -7,7 +7,7 @@ #include #include -#include "kmm/utils/macros.hpp" +#include "kmm/core/macros.hpp" namespace kmm { @@ -46,25 +46,34 @@ using gpu_mem_pool_attr_t = cudaMemPoolAttr; #define G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR +#define GPU_ERROR_PEER_ACCESS_ALREADY_ENABLED CUDA_ERROR_PEER_ACCESS_ALREADY_ENABLED #define g_init cuInit #define g_device_get_count cuDeviceGetCount #define g_device_get cuDeviceGet #define g_device_get_name cuDeviceGetName +#define g_device_total_mem cuDeviceTotalMem #define g_device_get_attribute cuDeviceGetAttribute +#define g_device_can_access_peer cuDeviceCanAccessPeer #define g_ctx_get_device cuCtxGetDevice +#define g_ctx_get_id cuCtxGetId #define g_ctx_create cuCtxCreate #define g_ctx_destroy cuCtxDestroy #define g_device_primary_ctx_retain cuDevicePrimaryCtxRetain #define g_device_primary_ctx_release cuDevicePrimaryCtxRelease #define g_ctx_push_current cuCtxPushCurrent #define g_ctx_pop_current cuCtxPopCurrent +#define g_ctx_enable_peer_access cuCtxEnablePeerAccess // Stream Management Constants & Functions #define G_STREAM_NON_BLOCKING CU_STREAM_NON_BLOCKING #define g_ctx_get_stream_priority_range cuCtxGetStreamPriorityRange +#define g_stream_create cuStreamCreate #define g_stream_create_with_priority cuStreamCreateWithPriority +#define g_stream_get_device cuStreamGetDevice +#define g_stream_get_ctx cuStreamGetCtx +#define g_stream_get_id cuStreamGetId #define g_stream_query cuStreamQuery #define g_stream_synchronize cuStreamSynchronize #define g_stream_destroy cuStreamDestroy @@ -94,16 +103,20 @@ using gpu_mem_pool_attr_t = cudaMemPoolAttr; #define G_MEMORYTYPE_DEVICE CU_MEMORYTYPE_DEVICE #define G_POINTER_ATTRIBUTE_MEMORY_TYPE CU_POINTER_ATTRIBUTE_MEMORY_TYPE #define G_POINTER_ATTRIBUTE_DEVICE_ORDINAL CU_POINTER_ATTRIBUTE_DEVICE_ORDINAL - -#define g_mem_get_info cuMemGetInfo -#define gpu_mem_get_info gpuMemGetInfo -#define g_mem_alloc cuMemAlloc -#define g_mem_free cuMemFree -#define gpu_malloc cudaMalloc -#define gpu_free cudaFree -#define g_mem_host_alloc cuMemHostAlloc -#define g_mem_free_host cuMemFreeHost -#define g_pointer_get_attribute cuPointerGetAttribute +#define G_MEM_ATTACH_GLOBAL CU_MEM_ATTACH_GLOBAL + +#define g_mem_get_info cuMemGetInfo +#define gpu_mem_get_info gpuMemGetInfo +#define g_mem_alloc cuMemAlloc +#define g_mem_free cuMemFree +#define gpu_malloc cudaMalloc +#define gpu_free cudaFree +#define g_mem_host_alloc cuMemHostAlloc +#define g_mem_free_host cuMemFreeHost +#define g_pointer_get_attribute cuPointerGetAttribute +#define g_mem_alloc_managed cuMemAllocManaged +#define g_mem_host_get_device_pointer cuMemHostGetDevicePointer +#define g_mem_prefetch_async cuMemPrefetchAsync // Memory Copy Operations @@ -122,16 +135,18 @@ using gpu_mem_pool_attr_t = cudaMemPoolAttr; #define g_memcpy_d_to_d_async cuMemcpyDtoDAsync #define g_memcpy_async cuMemcpyAsync -g_result_t g_memcpy_peer_async( - g_device_ptr_t, - g_context_t, - g_device_t, - g_device_ptr_t, - g_context_t, - g_device_t, - size_t, - g_stream_t -); +static inline g_result_t g_memcpy_peer_async( + g_device_ptr_t dstAddr, + g_context_t dstContext, + g_device_t dstDevice, + g_device_ptr_t srcAddr, + g_context_t srcContext, + g_device_t srcDevice, + size_t ByteCount, + g_stream_t hStream +) { + return cuMemcpyPeerAsync(dstAddr, dstContext, srcAddr, srcContext, ByteCount, hStream); +} // Memory Fill Operations @@ -145,6 +160,10 @@ g_result_t g_memcpy_peer_async( #define g_memcpy_2d cuMemcpy2D #define g_memcpy_2d_async cuMemcpy2DAsync +#define g_memset_d2d8_async cuMemsetD2D8Async +#define g_memset_d2d16_async cuMemsetD2D16Async +#define g_memset_d2d32_async cuMemsetD2D32Async + // Memory Pool Management Constants & Functions #define G_MEM_ALLOCATION_TYPE_PINNED CU_MEM_ALLOCATION_TYPE_PINNED diff --git a/include/kmm/backends/hip.hpp b/include/kmm/utils/backends/hip.hpp similarity index 54% rename from include/kmm/backends/hip.hpp rename to include/kmm/utils/backends/hip.hpp index 0c9436fc..c86026b5 100644 --- a/include/kmm/backends/hip.hpp +++ b/include/kmm/utils/backends/hip.hpp @@ -1,19 +1,26 @@ #pragma once #include -#include -#include +#if defined(__HIP__) + #include + #include +#endif #include #include -#include "kmm/utils/macros.hpp" +#include "kmm/core/macros.hpp" namespace kmm { // Types +#if defined(__HIP__) using half_type = __half; using bfloat16_type = __hip_bfloat16; +#else +using half_type = unsigned char; +using bfloat16_type = char; +#endif using g_result_t = hipError_t; using gpu_error_t = hipError_t; using g_device_t = hipDevice_t; @@ -45,12 +52,15 @@ using gpu_mem_pool_attr_t = hipMemPoolAttr; #define G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z hipDeviceAttributeMaxGridDimZ #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR hipDeviceAttributeComputeCapabilityMajor #define G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR hipDeviceAttributeComputeCapabilityMinor +#define GPU_ERROR_PEER_ACCESS_ALREADY_ENABLED hipErrorPeerAccessAlreadyEnabled #define g_init hipInit #define g_device_get_count hipGetDeviceCount #define g_device_get hipDeviceGet #define g_device_get_name hipDeviceGetName +#define g_device_total_mem hipDeviceTotalMem #define g_device_get_attribute hipDeviceGetAttribute +#define g_device_can_access_peer hipDeviceCanAccessPeer #define g_ctx_get_device hipCtxGetDevice #define g_ctx_create hipCtxCreate #define g_ctx_destroy hipCtxDestroy @@ -58,17 +68,59 @@ using gpu_mem_pool_attr_t = hipMemPoolAttr; #define g_device_primary_ctx_release hipDevicePrimaryCtxRelease #define g_ctx_push_current hipCtxPushCurrent #define g_ctx_pop_current hipCtxPopCurrent +#define g_ctx_enable_peer_access hipCtxEnablePeerAccess + +// Workaround because we rely on older versions of HIP +static inline g_result_t g_ctx_get_id(g_context_t context, unsigned long long* ctxId) { + *ctxId = reinterpret_cast(context); + return hipSuccess; +} // Stream Management Constants & Functions #define G_STREAM_NON_BLOCKING hipStreamNonBlocking #define g_ctx_get_stream_priority_range hipDeviceGetStreamPriorityRange +#define g_stream_create hipStreamCreateWithFlags #define g_stream_create_with_priority hipStreamCreateWithPriority +#define g_stream_get_device hipStreamGetDevice #define g_stream_query hipStreamQuery #define g_stream_synchronize hipStreamSynchronize #define g_stream_destroy hipStreamDestroy #define g_stream_wait_event hipStreamWaitEvent +static inline g_result_t g_stream_get_ctx(g_stream_t hStream, g_context_t* pctx) { + // HIP has no equivalent of `cuStreamGetCtx`. Instead, resolve the device + // the stream was created on and reuse its primary context, which is what + // `SystemInfo` already treats as "the" context for that device. + g_device_t device; + g_result_t result; + + result = g_stream_get_device(hStream, &device); + if (result != hipSuccess) { + return result; + } + + result = g_device_primary_ctx_retain(pctx, device); + if (result != hipSuccess) { + return result; + } + + return g_device_primary_ctx_release(device); +} + +// hipStreamGetId was only added in ROCm 7.1.0; on older ROCm (e.g. the 6.3.4 used in CI) there is +// no native stream-id query. Since g_stream_t (hipStream_t) is itself an opaque pointer that +// uniquely identifies the stream object for its lifetime, reuse the pointer value as the id -- +// this satisfies the only property callers (GPUStreamId) rely on: stable equality/hashing. +#if HIP_VERSION_MAJOR > 7 || (HIP_VERSION_MAJOR == 7 && HIP_VERSION_MINOR >= 1) + #define g_stream_get_id hipStreamGetId +#else +static inline g_result_t g_stream_get_id(g_stream_t hStream, unsigned long long* streamId) { + *streamId = reinterpret_cast(hStream); + return hipSuccess; +} +#endif + // Event Management Constants & Functions #define G_EVENT_WAIT_DEFAULT 0 @@ -91,18 +143,22 @@ using gpu_mem_pool_attr_t = hipMemPoolAttr; #define G_MEMHOSTALLOC_DEVICEMAP hipHostMallocMapped #define G_MEMORYTYPE_HOST hipMemoryTypeHost #define G_MEMORYTYPE_DEVICE hipMemoryTypeDevice -#define G_POINTER_ATTRIBUTE_DEVICE_ORDINAL HIP_POINTER_ATTRIBUTE_DEVICE_ORDINAL #define G_POINTER_ATTRIBUTE_MEMORY_TYPE HIP_POINTER_ATTRIBUTE_MEMORY_TYPE - -#define g_mem_get_info hipMemGetInfo -#define gpu_mem_get_info hipMemGetInfo -#define g_mem_alloc hipMalloc -#define g_mem_free hipFree -#define gpu_malloc hipMalloc -#define gpu_free hipFree -#define g_mem_host_alloc hipHostMalloc -#define g_mem_free_host hipHostFree -#define g_pointer_get_attribute hipPointerGetAttribute +#define G_POINTER_ATTRIBUTE_DEVICE_ORDINAL HIP_POINTER_ATTRIBUTE_DEVICE_ORDINAL +#define G_MEM_ATTACH_GLOBAL hipMemAttachGlobal + +#define g_mem_get_info hipMemGetInfo +#define gpu_mem_get_info hipMemGetInfo +#define g_mem_alloc hipMalloc +#define g_mem_free hipFree +#define gpu_malloc hipMalloc +#define gpu_free hipFree +#define g_mem_host_alloc hipHostMalloc +#define g_mem_free_host hipHostFree +#define g_pointer_get_attribute hipPointerGetAttribute +#define g_mem_alloc_managed hipMallocManaged +#define g_mem_host_get_device_pointer hipHostGetDevicePointer +#define g_mem_prefetch_async hipMemPrefetchAsync // Memory Copy Operations @@ -118,32 +174,44 @@ using gpu_mem_pool_attr_t = hipMemPoolAttr; #define g_memcpy_d_to_d hipMemcpyDtoD #define g_memcpy_d_to_d_async hipMemcpyDtoDAsync -g_result_t g_memcpy_async( +static inline g_result_t g_memcpy_async( g_device_ptr_t dst, g_device_ptr_t src, size_t ByteCount, g_stream_t hStream -); +) { + return hipMemcpyAsync(dst, src, ByteCount, hipMemcpyDefault, hStream); +} -g_result_t g_memcpy_h_to_d_async( +static inline g_result_t g_memcpy_h_to_d_async( g_device_ptr_t dstDevice, const void* srcHost, size_t ByteCount, g_stream_t hStream -); +) { + return hipMemcpyHtoDAsync(dstDevice, const_cast(srcHost), ByteCount, hStream); +} -g_result_t g_memcpy_h_to_d(g_device_ptr_t dstDevice, const void* srcHost, size_t ByteCount); - -g_result_t g_memcpy_peer_async( - g_device_ptr_t, - g_context_t, - g_device_t, - g_device_ptr_t, - g_context_t, - g_device_t, - size_t, - g_stream_t -); +static inline g_result_t g_memcpy_h_to_d( + g_device_ptr_t dstDevice, + const void* srcHost, + size_t ByteCount +) { + return hipMemcpyHtoD(dstDevice, const_cast(srcHost), ByteCount); +} + +static inline g_result_t g_memcpy_peer_async( + g_device_ptr_t dstAddr, + g_context_t dstContext, + g_device_t dstDevice, + g_device_ptr_t srcAddr, + g_context_t srcContext, + g_device_t srcDevice, + size_t ByteCount, + g_stream_t hStream +) { + return hipMemcpyPeerAsync(dstAddr, dstDevice, srcAddr, srcDevice, ByteCount, hStream); +} // Memory Fill Operations @@ -157,6 +225,94 @@ g_result_t g_memcpy_peer_async( #define g_memcpy_2d hipMemcpyParam2D #define g_memcpy_2d_async hipMemcpyParam2DAsync +// hipMemsetD2D{8,16,32}Async were only added in ROCm 7.1.0; on older ROCm (e.g. the +// 6.3.4 used in CI) emulate them with a per-row loop over the linear (1D) memsets, which do exist +// on ROCm 6. All calls go on the same stream, so this preserves the same stream-ordered semantics +// as a single native call would. +#if HIP_VERSION_MAJOR > 7 || (HIP_VERSION_MAJOR == 7 && HIP_VERSION_MINOR >= 1) + #define g_memset_d2d8_async hipMemsetD2D8Async + #define g_memset_d2d16_async hipMemsetD2D16Async + #define g_memset_d2d32_async hipMemsetD2D32Async +#else +static inline g_result_t g_memset_d2d8_async( + g_device_ptr_t dstDevice, + size_t dstPitch, + unsigned char uc, + size_t Width, + size_t Height, + g_stream_t hStream +) { + auto* base = reinterpret_cast(dstDevice); + + for (size_t row = 0; row < Height; row++) { + auto result = hipMemsetD8Async( + reinterpret_cast(base + row * dstPitch), + uc, + Width, + hStream + ); + + if (result != hipSuccess) { + return result; + } + } + + return hipSuccess; +} + +static inline g_result_t g_memset_d2d16_async( + g_device_ptr_t dstDevice, + size_t dstPitch, + unsigned short us, + size_t Width, + size_t Height, + g_stream_t hStream +) { + auto* base = reinterpret_cast(dstDevice); + + for (size_t row = 0; row < Height; row++) { + auto result = hipMemsetD16Async( + reinterpret_cast(base + row * dstPitch), + us, + Width, + hStream + ); + + if (result != hipSuccess) { + return result; + } + } + + return hipSuccess; +} + +static inline g_result_t g_memset_d2d32_async( + g_device_ptr_t dstDevice, + size_t dstPitch, + unsigned int ui, + size_t Width, + size_t Height, + g_stream_t hStream +) { + auto* base = reinterpret_cast(dstDevice); + + for (size_t row = 0; row < Height; row++) { + auto result = hipMemsetD32Async( + reinterpret_cast(base + row * dstPitch), + ui, + Width, + hStream + ); + + if (result != hipSuccess) { + return result; + } + } + + return hipSuccess; +} +#endif + // Memory Pool Management Constants & Functions #define G_MEM_ALLOCATION_TYPE_PINNED hipMemAllocationTypePinned @@ -180,7 +336,6 @@ g_result_t g_memcpy_peer_async( // Error Handling Constants & Functions -// HIP has no separate driver/runtime API split, so G_* and GPU_* map to the same symbols. #define G_SUCCESS hipSuccess #define G_ERROR_OUT_OF_MEMORY hipErrorOutOfMemory #define G_ERROR_UNKNOWN hipErrorUnknown @@ -213,6 +368,9 @@ using blas_handle_t = rocblas_handle; #define blas_set_stream rocblas_set_stream #define blas_destroy rocblas_destroy_handle #define blas_get_status_string rocblas_status_to_string -const char* blas_get_status_name(blas_status_t); + +static inline const char* blas_get_status_name(blas_status_t) { + return ""; +} } // namespace kmm diff --git a/include/kmm/utils/bounds.hpp b/include/kmm/utils/bounds.hpp deleted file mode 100644 index 30f1e801..00000000 --- a/include/kmm/utils/bounds.hpp +++ /dev/null @@ -1,334 +0,0 @@ -#pragma once - -#include "kmm/utils/checked_compare.hpp" -#include "kmm/utils/dim.hpp" -#include "kmm/utils/fixed_vector.hpp" -#include "kmm/utils/macros.hpp" -#include "kmm/utils/point.hpp" -#include "kmm/utils/range.hpp" - -namespace kmm { - -template -class Bounds: public fixed_vector, N> { - public: - using storage_type = fixed_vector, N>; - - KMM_HOST_DEVICE - explicit constexpr Bounds(const storage_type& storage) : storage_type(storage) {} - - KMM_HOST_DEVICE - Bounds() { - for (size_t i = 0; is_less(i, N); i++) { - (*this)[i] = Range(static_cast(0)); - } - } - - constexpr Bounds(const Bounds&) = default; - constexpr Bounds(Bounds&&) noexcept = default; - - Bounds& operator=(const Bounds&) = default; - Bounds& operator=(Bounds&&) noexcept = default; - - template - KMM_HOST_DEVICE constexpr Bounds(const Bounds& that) { - if (!that.template is_convertible_to()) { - throw_overflow_exception(); - } - - *this = Bounds::from(that); - } - - template::type> - KMM_HOST_DEVICE Bounds(Range first, Ts&&... args) : Bounds() { - (*this)[0] = first; - - size_t index = 0; - (((*this)[++index] = args), ...); - } - - KMM_HOST_DEVICE Bounds(const Dim& shape) { - *this = from_offset_size(Point::zero(), Dim::from(shape)); - } - - KMM_HOST_DEVICE static constexpr Bounds from_bounds( - const Point& begin, - const Point& end - ) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = {begin[i], end[i]}; - } - - return Bounds(result); - } - - KMM_HOST_DEVICE static constexpr Bounds from_offset_size( - const Point& offset, - const Dim& shape - ) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = Range(shape[i]).shift_by(offset[i]); - } - - return Bounds(result); - } - - KMM_HOST_DEVICE static constexpr Bounds empty() { - return Bounds(Dim::zero()); - } - - KMM_HOST_DEVICE static constexpr Bounds one() { - return Bounds(Dim::one()); - } - - template - KMM_HOST_DEVICE static constexpr Bounds from(const fixed_vector, M>& that) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = is_less(i, M) ? Range::from(that[i]) : Range(static_cast(1)); - } - - return Bounds(result); - } - - template - KMM_HOST_DEVICE bool is_convertible_to() const { - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - if (i < M) { - result &= (*this)[i].template is_convertible_to(); - } else { - result &= (*this)[i] == Range(static_cast(1)); - } - } - - return result; - } - - KMM_HOST_DEVICE - Range get_or_default(size_t i, Range default_value = static_cast(1)) const { - if constexpr (N > 0) { - if (KMM_LIKELY(i < N)) { - return (*this)[i]; - } - } - - return default_value; - } - - KMM_HOST_DEVICE - T begin(size_t axis) const { - return get_or_default(axis).begin; - } - - KMM_HOST_DEVICE - T end(size_t axis) const { - return get_or_default(axis).end; - } - - KMM_HOST_DEVICE - T size(size_t axis) const { - return get_or_default(axis).size(); - } - - KMM_HOST_DEVICE - Point begin() const { - Point result; - for (size_t axis = 0; is_less(axis, N); axis++) { - result[axis] = (*this)[axis].begin; - } - return result; - } - - KMM_HOST_DEVICE - Point end() const { - Point result; - for (size_t axis = 0; is_less(axis, N); axis++) { - result[axis] = (*this)[axis].end; - } - return result; - } - - KMM_HOST_DEVICE - Dim size() const { - Dim result; - for (size_t axis = 0; is_less(axis, N); axis++) { - result[axis] = (*this)[axis].size(); - } - return result; - } - - KMM_HOST_DEVICE - bool is_empty() const { - bool result = false; - - for (size_t i = 0; is_less(i, N); i++) { - result |= this->begin(i) >= this->end(i); - } - - return result; - } - - KMM_HOST_DEVICE - T volume() const { - T result = 1; - - for (size_t i = 0; is_less(i, N); i++) { - result *= this->end(i) - this->begin(i); - } - - return this->is_empty() ? T {0} : result; - } - - KMM_HOST_DEVICE - Bounds intersection(const Bounds& that) const { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i].begin = this->begin(i) >= that.begin(i) ? this->begin(i) : that.begin(i); - result[i].end = this->end(i) <= that.end(i) ? this->end(i) : that.end(i); - } - - return Bounds(result); - } - - KMM_HOST_DEVICE - Bounds unite(const Bounds& that) const { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i].begin = this->begin(i) < that.begin(i) ? this->begin(i) : that.begin(i); - result[i].end = this->end(i) > that.end(i) ? this->end(i) : that.end(i); - } - - return Bounds(result); - } - - KMM_HOST_DEVICE - bool overlaps(const Bounds& that) const { - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - result &= this->begin(i) < this->end(i) && that.begin(i) < that.end(i) && // - this->begin(i) < that.end(i) && that.begin(i) < this->end(i); - } - - return result; - } - - KMM_HOST_DEVICE - bool contains(const Bounds& that) const { - bool contains = true; - bool is_empty = false; - - for (size_t i = 0; is_less(i, N); i++) { - contains &= that.begin(i) >= this->begin(i); - contains &= that.end(i) <= this->end(i); - is_empty |= that.begin(i) >= that.end(i); - } - - return contains || is_empty; - } - - KMM_HOST_DEVICE - bool contains(const Point& that) const { - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - result &= (*this)[i].contains(that[i]); - } - - return result; - } - - template::type> - KMM_HOST_DEVICE bool contains(const T& first, Ts&&... rest) { - return contains(Point {first, rest...}); - } - - KMM_HOST_DEVICE - bool overlaps(const Dim& that) const { - return overlaps(Bounds {that}); - } - - KMM_HOST_DEVICE - bool contains(const Dim& that) const { - return contains(Bounds {that}); - } - - KMM_HOST_DEVICE - Bounds shift_by(const Point& offset) const { - storage_type result = *this; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = (*this)[i].shift_by(offset[i]); - } - - return Bounds(result); - } - - KMM_HOST_DEVICE - Bounds split_tail_along(size_t axis, const T& mid) { - if (is_less(axis, N)) { - auto result = *this; - result[axis] = (*this)[axis].split_tail(mid); - return result; - } else { - return Bounds::empty(); - } - } -}; - -template -Bounds(Ts&&...) -> Bounds; - -template -KMM_HOST_DEVICE Bounds concat(const Bounds& lhs, const Bounds& rhs) { - return Bounds { - concat((const fixed_vector, N>&)(lhs), (const fixed_vector, M>&)(rhs)) - }; -} - -template -KMM_HOST_DEVICE bool operator==(const Bounds& lhs, const Bounds& rhs) { - bool result = true; - - for (size_t i = 0; i < N || i < M; i++) { - result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); - } - - return result; -} - -template -KMM_HOST_DEVICE bool operator!=(const Bounds& lhs, const Bounds& rhs) { - return !(lhs == rhs); -} -} // namespace kmm - -#if !KMM_IS_RTC - #include - - #include "fmt/ostream.h" - - #include "kmm/utils/hash_utils.hpp" - -namespace kmm { -template -std::ostream& operator<<(std::ostream& stream, const Bounds& p) { - return stream << static_cast, N>&>(p); -} -} // namespace kmm - -template -struct fmt::formatter>: fmt::ostream_formatter {}; - -template -struct std::hash>: std::hash> {}; -#endif \ No newline at end of file diff --git a/include/kmm/utils/checked_compare.hpp b/include/kmm/utils/checked_compare.hpp deleted file mode 100644 index e8647da9..00000000 --- a/include/kmm/utils/checked_compare.hpp +++ /dev/null @@ -1,475 +0,0 @@ -#pragma once - -#include -#include - -#include "kmm/utils/macros.hpp" -#include "kmm/utils/panic.hpp" - -namespace kmm { - -namespace detail { - -enum class numeric_type_tag { // - signed_int, - unsigned_int, - floating_point, - other -}; - -template -struct numeric_type_traits { - static constexpr numeric_type_tag tag = numeric_type_tag::other; -}; - -template<> -struct numeric_type_traits { - static constexpr numeric_type_tag tag = numeric_type_tag::floating_point; -}; - -template<> -struct numeric_type_traits { - static constexpr numeric_type_tag tag = numeric_type_tag::floating_point; -}; - -#define KMM_DEFINE_INT_TRAITS(T) \ - template<> \ - struct numeric_type_traits { \ - static constexpr numeric_type_tag tag = numeric_type_tag::signed_int; \ - using unsigned_type = unsigned T; \ - \ - static constexpr signed T max_inclusive = \ - (signed T)(unsigned_type(~unsigned_type(0)) >> 1); \ - static constexpr signed T min_inclusive = ~max_inclusive; \ - \ - static constexpr float min_inclusive_float = min_inclusive; \ - static constexpr float max_exclusive_float = unsigned_type(max_inclusive) + 1; \ - }; \ - \ - template<> \ - struct numeric_type_traits { \ - static constexpr numeric_type_tag tag = numeric_type_tag::unsigned_int; \ - \ - static constexpr unsigned T min_inclusive = 0; \ - static constexpr unsigned T max_inclusive = \ - static_cast(~static_cast(0)); \ - \ - static constexpr float min_inclusive_float = min_inclusive; \ - static constexpr float max_exclusive_float = 2.0f * float(max_inclusive / 2 + 1); \ - }; - -KMM_DEFINE_INT_TRAITS(char) -KMM_DEFINE_INT_TRAITS(short) -KMM_DEFINE_INT_TRAITS(int) -KMM_DEFINE_INT_TRAITS(long) -KMM_DEFINE_INT_TRAITS(long long) - -template<> -struct numeric_type_traits { - static constexpr bool is_signed = static_cast(-1) < 0; - using unsigned_type = unsigned char; - - static constexpr numeric_type_tag tag = is_signed // - ? numeric_type_tag::signed_int - : numeric_type_tag::unsigned_int; - - static constexpr char max_inclusive = is_signed // - ? numeric_type_traits::max_inclusive - : numeric_type_traits::min_inclusive; - - static constexpr char min_inclusive = is_signed // - ? numeric_type_traits::min_inclusive - : numeric_type_traits::min_inclusive; - - static constexpr float min_inclusive_float = float(min_inclusive); - static constexpr float max_exclusive_float = float(max_inclusive) + 1.0F; -}; - -template< - typename L, - typename R, - numeric_type_tag = numeric_type_traits::tag, - numeric_type_tag = numeric_type_traits::tag> -struct checked_compare_impl; - -template -struct checked_compare_impl { - KMM_HOST_DEVICE - static constexpr bool is_equal(const T& left, const T& right) { - return left == right; - } - - KMM_HOST_DEVICE - static constexpr bool is_less(const T& left, const T& right) { - return left < right; - } -}; - -template -struct checked_compare_impl { - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return left == right; - } - - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - return left < right; - } -}; - -template -struct checked_compare_impl { - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return left == right; - } - - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - return left < right; - } -}; - -template -struct checked_compare_impl { - using UR = typename numeric_type_traits::unsigned_type; - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return right >= static_cast(0) && left == static_cast(right); - } - - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - return right >= static_cast(0) && left < static_cast(right); - } -}; - -template -struct checked_compare_impl { - using UL = typename numeric_type_traits::unsigned_type; - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return left >= static_cast(0) && static_cast
          (left) == right; - } - - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - return left < static_cast(0) || static_cast
            (left) < right; - } -}; - -template -struct checked_compare_impl< - L, - R, - numeric_type_tag::floating_point, - numeric_type_tag::floating_point> { - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return left == right; - } - - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - return left < right; - } -}; - -template -struct checked_compare_impl { - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - if (floor(left) < numeric_type_traits::min_inclusive_float) { - return true; - } - - if (floor(left) >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return static_cast(floor(left)) < right; - } - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - if (floor(left) != left) { - return false; - } - - if (left < numeric_type_traits::min_inclusive_float) { - return false; - } - - if (left >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return static_cast(left) == right; - } -}; - -template -struct checked_compare_impl { - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - if (ceil(right) < numeric_type_traits::min_inclusive_float) { - return false; - } - - if (ceil(right) >= numeric_type_traits::max_exclusive_float) { - return true; - } - - return left < static_cast(ceil(right)); - } - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return checked_compare_impl::is_equal(right, left); - } -}; - -template -struct checked_compare_impl< - L, - R, - numeric_type_tag::floating_point, - numeric_type_tag::unsigned_int> { - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - if (floor(left) < numeric_type_traits::min_inclusive_float) { - return true; - } - - if (floor(left) >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return static_cast(floor(left)) < right; - } - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - if (floor(left) != left) { - return false; - } - - if (left < numeric_type_traits::min_inclusive_float) { - return false; - } - - if (left >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return static_cast(left) == right; - } -}; - -template -struct checked_compare_impl< - L, - R, - numeric_type_tag::unsigned_int, - numeric_type_tag::floating_point> { - KMM_HOST_DEVICE - static constexpr bool is_less(L left, R right) { - if (ceil(right) < numeric_type_traits::min_inclusive_float) { - return false; - } - - if (ceil(right) >= numeric_type_traits::max_exclusive_float) { - return true; - } - - return left < static_cast(ceil(right)); - } - - KMM_HOST_DEVICE - static constexpr bool is_equal(L left, R right) { - return checked_compare_impl::is_equal(right, left); - } -}; - -template< - typename I, - typename O, - numeric_type_tag = numeric_type_traits::tag, - numeric_type_tag = numeric_type_traits::tag> -struct checked_convert_impl; - -template -struct checked_convert_impl { - static bool check(const T& input) { - return true; - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - return static_cast(input) == input; - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - return static_cast(input) == input; - } -}; - -template -struct checked_convert_impl< - I, - O, - numeric_type_tag::floating_point, - numeric_type_tag::floating_point> { - KMM_HOST_DEVICE - static bool check(const I& input) { - return (input != input) || input == static_cast(input); - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - return !checked_compare_impl::is_less(numeric_type_traits::max_inclusive, input); - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - return input >= static_cast(0) && // - !checked_compare_impl::is_less(numeric_type_traits::max_inclusive, input); - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - if (trunc(input) != input) { - return false; - } - - if (input < numeric_type_traits::min_inclusive_float) { - return false; - } - - if (input >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return true; - } -}; - -template -struct checked_convert_impl< - I, - O, - numeric_type_tag::floating_point, - numeric_type_tag::unsigned_int> { - static bool check(const I& input) { - if (trunc(input) != input) { - return false; - } - - if (input < static_cast(0)) { - return false; - } - - if (input >= numeric_type_traits::max_exclusive_float) { - return false; - } - - return true; - } -}; - -template -struct checked_convert_impl { - KMM_HOST_DEVICE - static bool check(const I& input) { - return checked_compare_impl::is_equal(input, static_cast(input)); - } -}; - -template -struct checked_convert_impl< - I, - O, - numeric_type_tag::unsigned_int, - numeric_type_tag::floating_point> { - KMM_HOST_DEVICE - static bool check(const I& input) { - return checked_compare_impl::is_equal(input, static_cast(input)); - } -}; - -} // namespace detail - -#if KMM_IS_DEVICE -// on the GPU, we just panic immediately -KMM_DEVICE void throw_overflow_exception() { - KMM_PANIC("overflow occurred in operation"); -} -#else -// on the host, we can throw an exception -[[noreturn]] void throw_overflow_exception(); -#endif - -template -KMM_HOST_DEVICE constexpr bool is_less(const L& left, const R& right) { - return detail::checked_compare_impl::is_less(left, right); -} - -template -KMM_HOST_DEVICE constexpr bool is_equal(const L& left, const R& right) { - return detail::checked_compare_impl::is_equal(left, right); -} - -template -KMM_HOST_DEVICE constexpr bool is_less_equal(const L& left, const R& right) { - return is_less(left, right) || is_equal(left, right); -} - -template -KMM_HOST_DEVICE constexpr bool is_greater(const L& left, const R& right) { - return is_less(right, left); -} - -template -KMM_HOST_DEVICE constexpr bool is_greater_equal(const L& left, const R& right) { - return is_less_equal(right, left); -} - -template -KMM_HOST_DEVICE bool is_convertible(const T& input) { - return detail::checked_convert_impl::check(input); -} - -template -KMM_HOST_DEVICE constexpr bool in_range(const T& input, const U& length) { - return !is_less(input, 0) && is_less(input, length) && is_convertible(input); -} - -template -KMM_HOST_DEVICE constexpr U checked_cast(const T& input) { - if (!is_convertible(input)) { - throw_overflow_exception(); - } - - return static_cast(input); -} - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/utils/checked_math.hpp b/include/kmm/utils/checked_math.hpp deleted file mode 100644 index 7d1a4452..00000000 --- a/include/kmm/utils/checked_math.hpp +++ /dev/null @@ -1,174 +0,0 @@ -#pragma once - -#include "checked_compare.hpp" -#include "macros.hpp" - -namespace kmm { - -namespace detail { - -template -struct checked_arithmetic_impl; - -#define KMM_IMPL_CHECKED_ARITHMETIC(T, ADD_FUN, SUB_FUN, MUL_FUN) \ - template<> \ - struct checked_arithmetic_impl { \ - static bool add(T lhs, T rhs, T* result) { \ - return ADD_FUN(lhs, rhs, result) == false; \ - } \ - \ - static bool sub(T lhs, T rhs, T* result) { \ - return SUB_FUN(lhs, rhs, result) == false; \ - } \ - \ - static bool mul(T lhs, T rhs, T* result) { \ - return MUL_FUN(lhs, rhs, result) == false; \ - } \ - }; - -KMM_IMPL_CHECKED_ARITHMETIC( - signed int, - __builtin_sadd_overflow, - __builtin_ssub_overflow, - __builtin_smul_overflow -) - -KMM_IMPL_CHECKED_ARITHMETIC( - signed long, - __builtin_saddl_overflow, - __builtin_ssubl_overflow, - __builtin_smull_overflow -) - -KMM_IMPL_CHECKED_ARITHMETIC( - signed long long, - __builtin_saddll_overflow, - __builtin_ssubll_overflow, - __builtin_smulll_overflow -) - -KMM_IMPL_CHECKED_ARITHMETIC( - unsigned int, - __builtin_uadd_overflow, - __builtin_usub_overflow, - __builtin_umul_overflow -) - -KMM_IMPL_CHECKED_ARITHMETIC( - unsigned long, - __builtin_uaddl_overflow, - __builtin_usubl_overflow, - __builtin_umull_overflow -) - -KMM_IMPL_CHECKED_ARITHMETIC( - unsigned long long, - __builtin_uaddll_overflow, - __builtin_usubll_overflow, - __builtin_umulll_overflow -) - -#define KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(T, R) \ - template<> \ - struct checked_arithmetic_impl { \ - static bool add(T lhs, T rhs, T* result) { \ - R temp = static_cast(lhs) + static_cast(rhs); \ - *result = static_cast(temp); \ - return detail::checked_convert_impl::check(temp); \ - } \ - \ - static bool sub(T lhs, T rhs, T* result) { \ - R temp = static_cast(lhs) - static_cast(rhs); \ - *result = static_cast(temp); \ - return detail::checked_convert_impl::check(temp); \ - } \ - \ - static bool mul(T lhs, T rhs, T* result) { \ - R temp = static_cast(lhs) * static_cast(rhs); \ - *result = static_cast(temp); \ - return detail::checked_convert_impl::check(temp); \ - } \ - }; - -KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(signed short, signed int) -KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(unsigned short, signed int) - -KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(signed char, signed int) -KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(unsigned char, signed int) -KMM_IMPL_CHECKED_ARITHMETIC_FORWARD(char, signed int) - -} // namespace detail - -template -KMM_HOST_DEVICE T checked_add(const T& left, const T& right) { - T output; - - if (!detail::checked_arithmetic_impl::add(left, right, &output)) { - throw_overflow_exception(); - } - - return output; -} - -template -KMM_HOST_DEVICE T checked_sub(const T& left, const T& right) { - T output; - - if (!detail::checked_arithmetic_impl::sub(left, right, &output)) { - throw_overflow_exception(); - } - - return output; -} - -template -KMM_HOST_DEVICE T checked_mul(const T& left, const T& right) { - T output; - - if (!detail::checked_arithmetic_impl::mul(left, right, &output)) { - throw_overflow_exception(); - } - - return output; -} - -template -KMM_HOST_DEVICE T checked_neg(const T& input) { - return checked_sub(static_cast(0), input); -} - -template -KMM_HOST_DEVICE U checked_sum(const T* begin, const T* end, U initial = T(0)) { - bool is_valid = true; - U accum = initial; - - for (const T* it = begin; it != end; it++) { - is_valid &= is_convertible(*it); - is_valid &= detail::checked_arithmetic_impl::add(static_cast(*it), accum, &accum); - } - - if (!is_valid) { - throw_overflow_exception(); - } - - return accum; -} - -template -KMM_HOST_DEVICE U checked_product(const T* begin, const T* end, U initial = T(1)) { - bool is_valid = true; - U accum = initial; - - for (const T* it = begin; it != end; it++) { - is_valid &= is_convertible(*it); - is_valid &= detail::checked_arithmetic_impl::mul(static_cast(*it), accum, &accum); - } - - if (!is_valid) { - throw_overflow_exception(); - } - - return accum; -} - -} // namespace kmm \ No newline at end of file diff --git a/include/kmm/utils/dim.hpp b/include/kmm/utils/dim.hpp deleted file mode 100644 index 06f1b73e..00000000 --- a/include/kmm/utils/dim.hpp +++ /dev/null @@ -1,230 +0,0 @@ -#pragma once - -#include "kmm/utils/checked_compare.hpp" -#include "kmm/utils/fixed_vector.hpp" -#include "kmm/utils/macros.hpp" -#include "kmm/utils/point.hpp" - -namespace kmm { - -template -struct Dim: public fixed_vector { - public: - using storage_type = fixed_vector; - - KMM_HOST_DEVICE - explicit constexpr Dim(const storage_type& storage) : storage_type(storage) {} - - KMM_HOST_DEVICE - constexpr Dim() { - for (size_t i = 0; is_less(i, N); i++) { - (*this)[i] = static_cast(1); - } - } - - constexpr Dim(const Dim&) = default; - constexpr Dim(Dim&&) noexcept = default; - Dim& operator=(const Dim&) = default; - Dim& operator=(Dim&&) noexcept = default; - - template::type> - KMM_HOST_DEVICE Dim(T first, Ts&&... args) : Dim() { - (*this)[0] = first; - - size_t index = 0; - (((*this)[++index] = args), ...); - } - - template - KMM_HOST_DEVICE constexpr Dim(const Dim& that) { - if (!that.template is_convertible_to()) { - throw_overflow_exception(); - } - - *this = Dim::from(that); - } - - template - KMM_HOST_DEVICE static constexpr Dim from(const fixed_vector& that) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = is_less(i, M) ? static_cast(that[i]) : static_cast(1); - } - - return Dim(result); - } - - KMM_HOST_DEVICE - static Dim fill(T value) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = value; - } - - return Dim(result); - } - - KMM_HOST_DEVICE - static Dim one() { - return fill(static_cast(1)); - } - - KMM_HOST_DEVICE - static Dim zero() { - return fill(static_cast(0)); - } - - KMM_HOST_DEVICE - T get_or_default(size_t i, T default_value = static_cast(1)) const { - if constexpr (N > 0) { - if (KMM_LIKELY(is_less(i, N))) { - return (*this)[i]; - } - } - - return default_value; - } - - template - KMM_HOST_DEVICE bool is_convertible_to() const { - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - if (i < M) { - result &= is_convertible((*this)[i]); - } else { - result &= is_equal((*this)[i], static_cast(1)); - } - } - - return result; - } - - KMM_HOST_DEVICE - bool is_empty() const { - bool result = false; - - for (size_t i = 0; is_less(i, N); i++) { - result |= !(static_cast(0) < (*this)[i]); - } - - return result; - } - - KMM_HOST_DEVICE - T volume() const { - T result = static_cast(1); - - if constexpr (N >= 1) { - result = (*this)[0]; - - for (size_t i = 1; is_less(i, N); i++) { - result *= (*this)[i]; - } - } - - return is_empty() ? static_cast(0) : result; - } - - template - KMM_HOST_DEVICE bool contains(const Point& p) const { - bool result = true; - - for (size_t i = 0; is_less(i, N) && is_less(i, M); i++) { - result &= !is_less(p[i], static_cast(0)) && is_less(p[i], (*this)[i]); - } - - if constexpr (N < M) { - for (size_t i = N; is_less(i, M); i++) { - result &= is_equal(p[i], static_cast(0)); - } - } - - if constexpr (N > M) { - for (size_t i = M; is_less(i, N); i++) { - result &= is_less(static_cast(0), (*this)[i]); - } - } - - return result; - } -}; - -template -Dim(Ts&&...) -> Dim; - -template -KMM_HOST_DEVICE Dim concat(const Dim& lhs, const Dim& rhs) { - return Dim {concat((const fixed_vector&)(lhs), (const fixed_vector&)(rhs)) - }; -} - -template -KMM_HOST_DEVICE bool operator==(const Dim& lhs, const Dim& rhs) { - bool result = true; - - for (size_t i = 0; is_less(i, N) || is_less(i, M); i++) { - result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); - } - - return result; -} - -template -KMM_HOST_DEVICE bool operator!=(const Dim& lhs, const Dim& rhs) { - return !(lhs == rhs); -} - -namespace detail { -// Specialize comparison between Dim and T -template -struct checked_compare_impl, Ltag, Rtag>: checked_compare_impl {}; - -template -struct checked_compare_impl, R, Ltag, Rtag>: checked_compare_impl {}; - -template -struct checked_compare_impl, Dim<1, R>, Ltag, Rtag>: checked_compare_impl {}; - -template -struct checked_compare_impl, Dim<1, T>, numeric_type_tag::other, numeric_type_tag::other>: - checked_compare_impl {}; - -template -struct checked_convert_impl, Ltag, Rtag>: checked_convert_impl {}; - -template -struct checked_convert_impl, R, Ltag, Rtag>: checked_convert_impl {}; - -template -struct checked_convert_impl, Dim<1, R>, Ltag, Rtag>: checked_convert_impl {}; - -template -struct checked_convert_impl, Dim<1, T>, numeric_type_tag::other, numeric_type_tag::other>: - checked_convert_impl {}; -} // namespace detail - -} // namespace kmm - -#if !KMM_IS_RTC - #include - - #include "fmt/ostream.h" - - #include "kmm/utils/hash_utils.hpp" - -namespace kmm { -template -std::ostream& operator<<(std::ostream& stream, const Dim& p) { - return stream << static_cast&>(p); -} -} // namespace kmm - -template -struct fmt::formatter>: fmt::ostream_formatter {}; - -template -struct std::hash>: std::hash> {}; -#endif \ No newline at end of file diff --git a/include/kmm/utils/fixed_vector.hpp b/include/kmm/utils/fixed_vector.hpp deleted file mode 100644 index 5f178fef..00000000 --- a/include/kmm/utils/fixed_vector.hpp +++ /dev/null @@ -1,280 +0,0 @@ -#pragma once - -#include "checked_compare.hpp" -#include "panic.hpp" - -namespace kmm { - -#define KMM_PANIC_OUT_OF_BOUNDS() KMM_PANIC("access out of bounds") - -template -static constexpr size_t compute_fixed_vector_alignment(size_t num) { - constexpr size_t max_align = 16; - size_t align = alignof(T); - while (align < max_align && align < num * sizeof(T)) { - align *= 2; - } - return align; -} - -template -struct alignas(compute_fixed_vector_alignment(N)) fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - if (i >= N) { - KMM_PANIC_OUT_OF_BOUNDS(); - } - - return __internal_data[i]; - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - if (i >= N) { - KMM_PANIC_OUT_OF_BOUNDS(); - } - - return __internal_data[i]; - } - - T __internal_data[N] {}; -}; - -template -struct fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - KMM_PANIC_OUT_OF_BOUNDS(); - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - KMM_PANIC_OUT_OF_BOUNDS(); - } -}; - -template -struct alignas(compute_fixed_vector_alignment(1)) fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - switch (i) { - case 0: - return x; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - switch (i) { - case 0: - return x; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - KMM_HOST_DEVICE - operator T() const { - return x; - } - - KMM_HOST_DEVICE - fixed_vector& operator=(T value) { - x = value; - return *this; - } - - T x {}; -}; - -template -struct alignas(compute_fixed_vector_alignment(2)) fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - switch (i) { - case 0: - return x; - case 1: - return y; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - switch (i) { - case 0: - return x; - case 1: - return y; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - T x {}; - T y {}; -}; - -template -struct alignas(compute_fixed_vector_alignment(3)) fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - switch (i) { - case 0: - return x; - case 1: - return y; - case 2: - return z; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - switch (i) { - case 0: - return x; - case 1: - return y; - case 2: - return z; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - T x {}; - T y {}; - T z {}; -}; - -template -struct alignas(compute_fixed_vector_alignment(4)) fixed_vector { - KMM_HOST_DEVICE - T& operator[](size_t i) { - switch (i) { - case 0: - return x; - case 1: - return y; - case 2: - return z; - case 3: - return w; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - KMM_HOST_DEVICE - const T& operator[](size_t i) const { - switch (i) { - case 0: - return x; - case 1: - return y; - case 2: - return z; - case 3: - return w; - default: - KMM_PANIC_OUT_OF_BOUNDS(); - } - } - - T x {}; - T y {}; - T z {}; - T w {}; -}; - -template -KMM_HOST_DEVICE bool operator==(const fixed_vector& lhs, const fixed_vector& rhs) { - if (N != M) { - return false; - } - - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - result &= is_equal(lhs[i], rhs[i]); - } - - return result; -} - -template -KMM_HOST_DEVICE bool operator!=(const fixed_vector& lhs, const fixed_vector& rhs) { - return !(lhs == rhs); -} - -template -KMM_HOST_DEVICE fixed_vector concat( - const fixed_vector& lhs, - const fixed_vector& rhs -) { - fixed_vector result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = lhs[i]; - } - - for (size_t i = 0; is_less(i, M); i++) { - result[i + N] = rhs[i]; - } - - return result; -} -} // namespace kmm - -#if !KMM_IS_RTC - #include - - #include "fmt/ostream.h" - - #include "kmm/utils/hash_utils.hpp" - -namespace kmm { - -template -std::ostream& operator<<(std::ostream& stream, const fixed_vector& p) { - stream << "{"; - for (size_t i = 0; is_less(i, N); i++) { - if (i != 0) { - stream << ", "; - } - - stream << p[i]; - } - - return stream << "}"; -} -} // namespace kmm - -template -struct fmt::formatter>: fmt::ostream_formatter {}; - -template -struct std::hash> { - size_t operator()(const kmm::fixed_vector& p) const { - size_t result = 0; - for (size_t i = 0; i < N; i++) { - kmm::hash_combine(result, p[i]); - } - return result; - } -}; - -template -struct std::hash> { - size_t operator()(const kmm::fixed_vector& p) const { - return 0; - } -}; -#endif \ No newline at end of file diff --git a/include/kmm/utils/function_ref.hpp b/include/kmm/utils/function_ref.hpp new file mode 100644 index 00000000..4e6f45df --- /dev/null +++ b/include/kmm/utils/function_ref.hpp @@ -0,0 +1,52 @@ +#pragma once + +#include +#include + +namespace kmm { + +/// \addtogroup utility +/// @{ + +template +class function_ref; + +/** + * A lightweight, non-owning reference to a callable, similar to `std::function_ref`. + */ +template +class function_ref { + public: + function_ref() noexcept = default; + function_ref(decltype(nullptr)) noexcept {} + + template< + typename Fn, + typename = std::enable_if_t< + !std::is_same_v, function_ref> + && std::is_invocable_r_v>> + function_ref(Fn&& callable) noexcept : + m_ptr(reinterpret_cast(std::addressof(callable))), + m_invoke(&invoke>) {} + + R operator()(Args... args) const { + return m_invoke(m_ptr, std::forward(args)...); + } + + explicit operator bool() const noexcept { + return m_ptr != nullptr; + } + + private: + template + static R invoke(void* ptr, Args... args) { + return (*reinterpret_cast(ptr))(std::forward(args)...); + } + + void* m_ptr = nullptr; + R (*m_invoke)(void*, Args...) = nullptr; +}; + +/// @} + +} // namespace kmm diff --git a/include/kmm/utils/geometry.hpp b/include/kmm/utils/geometry.hpp deleted file mode 100644 index 30dea48c..00000000 --- a/include/kmm/utils/geometry.hpp +++ /dev/null @@ -1,4 +0,0 @@ -#include "kmm/utils/bounds.hpp" -#include "kmm/utils/dim.hpp" -#include "kmm/utils/point.hpp" -#include "kmm/utils/range.hpp" \ No newline at end of file diff --git a/include/kmm/utils/gpu_utils.hpp b/include/kmm/utils/gpu_utils.hpp index d46c6977..662a89b8 100644 --- a/include/kmm/utils/gpu_utils.hpp +++ b/include/kmm/utils/gpu_utils.hpp @@ -1,26 +1,31 @@ #pragma once +#include +#include #include -#include +#include #include -#include +#include -#include "kmm/core/backends.hpp" -#include "kmm/utils/macros.hpp" +#include "fmt/ostream.h" + +#include "kmm/core/macros.hpp" +#include "kmm/utils/backends.hpp" #define KMM_GPU_CHECK(...) \ do { \ auto __code = (__VA_ARGS__); \ - if (__code != decltype(__code)(0)) { \ + if (KMM_UNLIKELY(__code != decltype(__code)(0))) { \ ::kmm::gpu_throw_exception(__code, __FILE__, __LINE__, #__VA_ARGS__); \ } \ - } while (0); + } while (0) namespace kmm { +/// \addtogroup utility +/// @{ + void gpu_throw_exception(g_result_t result, const char* file, int line, const char* expression); -void gpu_throw_exception(gpu_error_t result, const char* file, int line, const char* expression); -void gpu_throw_exception(blas_status_t result, const char* file, int line, const char* expression); class GPUException: public std::exception { public: @@ -34,76 +39,163 @@ class GPUException: public std::exception { std::string m_message; }; -class GPUDriverException: public GPUException { - public: - GPUDriverException(const std::string& message, g_result_t result); - GPUDriverException(const char* message, g_result_t result) : - GPUDriverException(std::string(message), result) {} - g_result_t status; -}; +class GPUContextGuard { + KMM_NOT_COPYABLE_OR_MOVABLE(GPUContextGuard) -class GPURuntimeException: public GPUException { public: - GPURuntimeException(const std::string& message, gpu_error_t result); - gpu_error_t status; + GPUContextGuard(g_context_t context); + ~GPUContextGuard(); + + private: + g_context_t m_context; }; -class BlasException: public GPUException { - public: - BlasException(const std::string& message, blas_status_t result); - blas_status_t status; +g_context_t context_from_stream(g_stream_t stream); + +struct GPUContextId { + GPUContextId(g_context_t context); + + bool operator==(const GPUContextId& that) const noexcept { + return m_id == that.m_id; + } + + bool operator!=(const GPUContextId& that) const noexcept { + return !(*this == that); + } + + unsigned long long get() const noexcept { + return m_id; + } + + friend std::ostream& operator<<(std::ostream&, const GPUContextId& self); + + private: + unsigned long long m_id; }; -/** - * Returns the available devices as a list of `device`s. - */ -std::vector get_gpu_devices(); +class GPUStreamRef; +class GPUStreamOwner; + +struct GPUStreamId { + GPUStreamId(g_stream_t stream); + GPUStreamId(g_stream_t stream, g_context_t context); + GPUStreamId(const GPUStreamRef& stream); + GPUStreamId(const GPUStreamOwner& stream); + + unsigned long long get() const noexcept { + return m_id; + } + + const GPUContextId& context() const noexcept { + return m_context_id; + } + + bool operator==(const GPUStreamId& that) const noexcept { + return m_id == that.m_id && m_context_id == that.m_context_id; + } + + bool operator!=(const GPUStreamId& that) const noexcept { + return !(*this == that); + } -/** - * If the given address points to memory allocation that has been allocated on a GPU, then - * this function returns the device ordinal as a `device`. If the address points ot an invalid - * memory location or a non-GPU buffer, then it returns `std::nullopt`. - */ -std::optional get_gpu_device_by_address(const void* address); + friend std::ostream& operator<<(std::ostream&, const GPUStreamId& self); -class GPUContextHandle { - GPUContextHandle() = delete; - GPUContextHandle(g_context_t context, std::shared_ptr lifetime); + private: + GPUContextId m_context_id; + unsigned long long m_id; +}; +class GPUStreamRef { public: - static GPUContextHandle create_context_for_device(g_device_t device); - static GPUContextHandle retain_primary_context_for_device(g_device_t device); + GPUStreamRef(g_stream_t); - operator g_context_t() const { + g_context_t context() const noexcept { return m_context; } + g_stream_t stream() const noexcept { + return m_stream; + } + + GPUStreamId stream_id() const noexcept { + return m_stream_id; + } + + operator g_stream_t() const noexcept { + return m_stream; + } + + friend std::ostream& operator<<(std::ostream&, const GPUStreamRef& self); + private: - g_context_t m_context; - std::shared_ptr m_lifetime; + g_context_t m_context = nullptr; + g_stream_t m_stream = nullptr; + GPUStreamId m_stream_id; }; -inline bool operator==(const GPUContextHandle& lhs, const GPUContextHandle& rhs) { - return g_context_t(lhs) == g_context_t(rhs); -} +class GPUStreamOwner { + public: + GPUStreamOwner(const GPUStreamOwner&) = delete; + GPUStreamOwner& operator=(const GPUStreamOwner&) = delete; + + explicit GPUStreamOwner(g_context_t context, unsigned int flags = G_STREAM_NON_BLOCKING); + ~GPUStreamOwner(); -inline bool operator!=(const GPUContextHandle& lhs, const GPUContextHandle& rhs) { - return !(lhs == rhs); -} + GPUStreamOwner(GPUStreamOwner&& that) noexcept : m_stream(that.m_stream) { + that.m_stream = nullptr; + } -class GPUContextGuard { - KMM_NOT_COPYABLE_OR_MOVABLE(GPUContextGuard) + GPUStreamOwner& operator=(GPUStreamOwner&& that) noexcept { + std::swap(m_stream, that.m_stream); + return *this; + } - public: - GPUContextGuard(GPUContextHandle context); - ~GPUContextGuard(); + g_stream_t get() const noexcept { + return m_stream; + } + + operator g_stream_t() const noexcept { + return m_stream; + } + + operator GPUStreamRef() const noexcept { + return m_stream; + } + + friend std::ostream& operator<<(std::ostream&, const GPUStreamOwner& self); private: - GPUContextHandle m_context; + void destroy() noexcept; + + g_stream_t m_stream = nullptr; }; -inline g_device_ptr_t gpu_deviceptr_offset(g_device_ptr_t ptr, size_t size) { - return reinterpret_cast(reinterpret_cast(ptr) + size); -} +/// @} } // namespace kmm + +template<> +struct std::hash { + size_t operator()(const kmm::GPUContextId& id) const noexcept { + return std::hash {}(id.get()); + } +}; + +template<> +struct std::hash { + size_t operator()(const kmm::GPUStreamId& id) const noexcept { + return std::hash {}(id.get()); + } +}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; diff --git a/include/kmm/utils/hash_utils.hpp b/include/kmm/utils/hash_utils.hpp index d8f56ca3..fe2ae6bb 100644 --- a/include/kmm/utils/hash_utils.hpp +++ b/include/kmm/utils/hash_utils.hpp @@ -1,13 +1,27 @@ #pragma once -#include -#include +#include "kmm/core/macros.hpp" + +// This header is only ever meant to be pulled in from behind another header's own +// `#if !KMM_IS_RTC` guard (see vec.hpp/point.hpp/range.hpp), but it is self-guarded here too so +// that including it directly can never break NVRTC device compilation. +#if !KMM_IS_RTC + #include + #include namespace kmm { +/// \addtogroup utility +/// @{ + template struct Hasher: std::hash {}; +template> +size_t hash_value(const T& v, H hasher = {}) { + return hasher(v); +} + template> void hash_combine(size_t& seed, const T& v, H hasher = {}) { seed ^= hasher(v) + 0x9e3779b9 + (seed << 6) + (seed >> 2); @@ -34,4 +48,7 @@ size_t hash_range(It begin, It end) { return seed; } -} // namespace kmm \ No newline at end of file +/// @} + +} // namespace kmm +#endif diff --git a/include/kmm/utils/integer_fun.hpp b/include/kmm/utils/integer_fun.hpp deleted file mode 100644 index e572406a..00000000 --- a/include/kmm/utils/integer_fun.hpp +++ /dev/null @@ -1,113 +0,0 @@ -#pragma once - -namespace kmm { - -namespace details { -template -static constexpr bool is_signed_integral = T(-1) < T(0); -} - -/** - * Divide `num` by `denom` and round the result down. - */ -template -constexpr T div_floor(T a, T b) { - constexpr T zero = static_cast(0); - T quotient = a / b; - - if constexpr (details::is_signed_integral) { - // Adjust the quotient if a and b have different signs - if (a % b != zero && ((a >= zero) ^ (b >= zero))) { - quotient -= 1; - } - } - - return quotient; -} - -/** - * Divide `num` by `denom` and round the result up. - */ -template -constexpr T div_ceil(T a, T b) { - constexpr T zero = static_cast(0); - T quotient = a / b; - - // Adjust the quotient - if (a % b != zero) { - if constexpr (details::is_signed_integral) { - // Adjust the quotient if both a and b have the same sign - if ((a >= zero) == (b >= zero)) { - quotient += 1; - } - } else { - quotient += 1; - } - } - - return quotient; -} - -/** - * Round `input` to the first multiple of `multiple`. - * - * In other words, returns the smallest value not less than `input` that is divisible by `multiple`. - */ -template -constexpr T round_up_to_multiple(T input, T multiple) { - constexpr T zero = static_cast(0); - - if constexpr (details::is_signed_integral) { - if (multiple < zero) { - multiple = -multiple; - } - } - - T remainder = input % multiple; - - if (remainder == zero) { - return input; - } else { - if constexpr (details::is_signed_integral) { - if (input < zero) { - return input - remainder; - } else { - return input - remainder + multiple; - } - } else { - return input - remainder + multiple; - } - } -} - -/** - * Return the smallest integer that is a power of two and is not less than `input`. - */ -template -constexpr T round_up_to_power_of_two(T input) { - if (input <= static_cast(0)) { - return static_cast(1); - } - - input -= static_cast(1); - for (decltype(sizeof(T)) i = 1; i < sizeof(T) * 8; i *= 2) { - input |= (input >> i); - } - - input += static_cast(1); - return input; -} - -/** - * Check if the given integer is a power of two. - */ -template -static bool is_power_of_two(T input) { - if (input <= static_cast(0)) { - return false; - } - - return (input & (input - 1)) == static_cast(0); -} - -} // namespace kmm diff --git a/include/kmm/utils/lru_cache.hpp b/include/kmm/utils/lru_cache.hpp new file mode 100644 index 00000000..961888b2 --- /dev/null +++ b/include/kmm/utils/lru_cache.hpp @@ -0,0 +1,116 @@ +#pragma once + +#include +#include +#include + +namespace kmm { + +/// \addtogroup utility +/// @{ + +/// A cache mapping keys to values that tracks usage order, so the caller can find and evict +/// the least-recently-used entry. `find`, `touch`, and `insert` all mark the given key as most +/// recently used. Eviction is not automatic; call `least_recently_used` and `remove` to evict. +template> +class lru_cache { + struct entry { + K key; + V value; + }; + + using list_type = std::list; + using map_type = ankerl::unordered_dense::map; + + public: + size_t size() const noexcept { + return m_map.size(); + } + + bool is_empty() const noexcept { + return m_map.empty(); + } + + bool contains(const K& key) const { + return m_map.find(key) != m_map.end(); + } + + /// Returns a pointer to the value associated with `key`, marking it as most recently used, + /// or `nullptr` if `key` is not present. + V* find(const K& key) { + auto it = m_map.find(key); + + if (it == m_map.end()) { + return nullptr; + } + + move_to_front(it->second); + return &it->second->value; + } + + const V* find(const K& key) const { + return const_cast(this)->find(key); + } + + /// Marks `key` as most recently used. Does nothing if `key` is not present. + void touch(const K& key) { + auto it = m_map.find(key); + + if (it != m_map.end()) { + move_to_front(it->second); + } + } + + /// Inserts or overwrites the value for `key`, marking it as most recently used. + void insert(K key, V value) { + auto it = m_map.find(key); + + if (it != m_map.end()) { + it->second->value = std::move(value); + move_to_front(it->second); + return; + } + + m_list.push_front(entry {key, std::move(value)}); + m_map.emplace(std::move(key), m_list.begin()); + } + + /// Removes `key` from the cache. Returns `true` if `key` was present. + bool remove(const K& key) { + auto it = m_map.find(key); + + if (it == m_map.end()) { + return false; + } + + m_list.erase(it->second); + m_map.erase(it); + return true; + } + + /// Returns the key of the least-recently-used entry, or `nullptr` if the cache is empty. + const K* least_recently_used() const noexcept { + if (m_list.empty()) { + return nullptr; + } + + return &m_list.back().key; + } + + void clear() noexcept { + m_list.clear(); + m_map.clear(); + } + + private: + void move_to_front(typename list_type::iterator it) { + m_list.splice(m_list.begin(), m_list, it); + } + + list_type m_list; + map_type m_map; +}; + +/// @} + +} // namespace kmm diff --git a/include/kmm/utils/notify.hpp b/include/kmm/utils/notify.hpp index 2c9df891..daba15be 100644 --- a/include/kmm/utils/notify.hpp +++ b/include/kmm/utils/notify.hpp @@ -1,82 +1,84 @@ #pragma once #include -#include -#include + +#include "kmm/utils/refcnt_ptr.hpp" namespace kmm { +/// \addtogroup utility +/// @{ + /** - * Interface to notify when an event has occurred + * Simple interface having a `notify` method to be called when an event happens. */ -class Notify { +class Notify: public reference_count { public: virtual ~Notify() noexcept = default; - - /** - * Called when an event is triggered. - */ virtual void notify() const noexcept = 0; }; /** - * Implementation of `Notify` that forwards the notify call to a function `F`. + * Wrapper around a `std::shared_ptr`. */ -template -class NotifyImpl: public Notify { +class NotifyHandle { public: - NotifyImpl(F fun = {}) : m_callback(std::move(fun)) {} + NotifyHandle() = default; + ~NotifyHandle(); - void notify() const noexcept final { - m_callback(); - } + /** Constructs an empty (null) handle. */ + NotifyHandle(decltype(nullptr)) {}; - private: - F m_callback; -}; + /** Constructs a handle from a shared pointer to a `Notify` instance. */ + NotifyHandle(refcnt_ptr m) : m_impl(std::move(m)) {} -/** - * Wrapper around a `shared_ptr` handle. - */ -class NotifyHandle { - public: - NotifyHandle() = default; - NotifyHandle(std::shared_ptr m); - NotifyHandle(std::unique_ptr m); + /** Constructs a handle by taking ownership of a `Notify` via unique pointer. */ + NotifyHandle(std::unique_ptr m) : m_impl(std::move(m)) {} template - NotifyHandle(std::shared_ptr m) : NotifyHandle(std::shared_ptr(m)) {} + NotifyHandle(refcnt_ptr m) : m_impl(std::move(m)) {} template - NotifyHandle(std::unique_ptr m) : NotifyHandle(std::shared_ptr(std::move(m))) {} + NotifyHandle(std::unique_ptr m) : m_impl(std::move(m)) {} - template::value, int>::type = 0> + /// Constructs a handle from any callable (e.g. lambda, functor, function pointer). + /// The callable is wrapped in an internal `Notify` implementation and stored via shared pointer. + template>>> NotifyHandle(F&& callback) : - NotifyHandle( - std::make_shared>::type>(std::forward(callback)) - ) {} + m_impl(make_refcnt>>(std::forward(callback))) {} - ~NotifyHandle(); - - /** - * If the underlying `Notify` object exists, this calls its `notify()` method. - */ + /// Calls `notify` on the inner notifier. Does nothing if the handle is empty. void notify() const noexcept; - /** - * Resets the managed `shared_ptr` to null, effectively clearing the - * notification handler. - */ + /// Resets the handle, releasing the reference to the inner notifier. void clear() noexcept; - /** - * This calls the `notify()` method on the underlying `Notify` object (if it exists), - * and then resets the managed `shared_ptr` to null. - */ + /// Notifies and then clears the handle. Equivalent to `notify(); clear();`. void notify_and_clear() noexcept; + /// Returns `true` if the handle holds a notifier, `false` if it is empty. + explicit operator bool() const noexcept { + return m_impl != nullptr; + } + private: - std::shared_ptr m_impl; + // Internal `Notify` implementation that wraps a callable `F`. + template + class Impl: public Notify { + public: + explicit Impl(F fun) : m_callback(std::move(fun)) {} + + void notify() const noexcept override { + m_callback(); + } + + private: + F m_callback; + }; + + refcnt_ptr m_impl; }; +/// @} + } // namespace kmm \ No newline at end of file diff --git a/include/kmm/utils/point.hpp b/include/kmm/utils/point.hpp deleted file mode 100644 index 226302eb..00000000 --- a/include/kmm/utils/point.hpp +++ /dev/null @@ -1,159 +0,0 @@ -#pragma once - -#include "kmm/utils/checked_compare.hpp" -#include "kmm/utils/fixed_vector.hpp" -#include "kmm/utils/macros.hpp" - -namespace kmm { - -using default_index_type = signed long int; // int64_t - -namespace detail { -template -struct enable_if {}; - -template -struct enable_if { - using type = T; -}; -} // namespace detail - -template -class Point: public fixed_vector { - public: - using storage_type = fixed_vector; - - KMM_HOST_DEVICE - explicit constexpr Point(const storage_type& storage) : storage_type(storage) {} - - KMM_HOST_DEVICE - constexpr Point() { - for (size_t i = 0; is_less(i, N); i++) { - (*this)[i] = T {}; - } - } - - template::type> - KMM_HOST_DEVICE Point(T first, Ts&&... args) : Point() { - (*this)[0] = first; - - size_t index = 0; - (((*this)[++index] = args), ...); - } - - template - KMM_HOST_DEVICE constexpr Point(const Point& that) { - if (!that.template is_convertible_to()) { - throw_overflow_exception(); - } - - *this = Point::from(that); - } - - template - KMM_HOST_DEVICE static constexpr Point from(const fixed_vector& that) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = is_less(i, M) ? static_cast(that[i]) : static_cast(0); - } - - return Point(result); - } - - KMM_HOST_DEVICE - static Point fill(T value) { - storage_type result; - - for (size_t i = 0; is_less(i, N); i++) { - result[i] = value; - } - - return Point(result); - } - - KMM_HOST_DEVICE - static Point one() { - return fill(static_cast(1)); - } - - KMM_HOST_DEVICE - static Point zero() { - return fill(static_cast(0)); - } - - template - KMM_HOST_DEVICE bool is_convertible_to() const { - bool result = true; - - for (size_t i = 0; is_less(i, N); i++) { - if (is_less(i, M)) { - result &= is_convertible((*this)[i]); - } else { - result &= is_equal((*this)[i], static_cast(0)); - } - } - - return result; - } - - KMM_HOST_DEVICE - T get_or_default(size_t i, T default_value = {}) const { - if constexpr (N > 0) { - if (KMM_LIKELY(is_less(i, N))) { - return (*this)[i]; - } - } - - return default_value; - } -}; - -template -Point(Ts&&...) -> Point; - -template -KMM_HOST_DEVICE Point concat(const Point& lhs, const Point& rhs) { - return Point { - concat((const fixed_vector&)(lhs), (const fixed_vector&)(rhs)) - }; -} - -template -KMM_HOST_DEVICE bool operator==(const Point& lhs, const Point& rhs) { - bool result = true; - - for (size_t i = 0; is_less(i, N) || is_less(i, M); i++) { - result &= is_equal(lhs.get_or_default(i), rhs.get_or_default(i)); - } - - return result; -} - -template -KMM_HOST_DEVICE bool operator!=(const Point& lhs, const Point& rhs) { - return !(lhs == rhs); -} - -} // namespace kmm - -#if !KMM_IS_RTC - #include - - #include "fmt/ostream.h" - - #include "kmm/utils/hash_utils.hpp" - -namespace kmm { -template -std::ostream& operator<<(std::ostream& stream, const Point& p) { - return stream << static_cast&>(p); -} -} // namespace kmm - -template -struct fmt::formatter>: fmt::ostream_formatter {}; - -template -struct std::hash>: std::hash> {}; -#endif \ No newline at end of file diff --git a/include/kmm/utils/poll.hpp b/include/kmm/utils/poll.hpp index 2804e7fe..a1b95b63 100644 --- a/include/kmm/utils/poll.hpp +++ b/include/kmm/utils/poll.hpp @@ -1,9 +1,12 @@ #pragma once -#include - namespace kmm { +/// \addtogroup utility +/// @{ + enum struct Poll { Ready, Pending }; +/// @} + } // namespace kmm \ No newline at end of file diff --git a/include/kmm/utils/range.hpp b/include/kmm/utils/range.hpp deleted file mode 100644 index f4318f3c..00000000 --- a/include/kmm/utils/range.hpp +++ /dev/null @@ -1,175 +0,0 @@ -#pragma once - -#include "checked_compare.hpp" - -namespace kmm { - -template -class Range { - public: - using value_type = T; - - constexpr Range(const Range&) = default; - constexpr Range(Range&&) = default; - - Range& operator=(const Range&) = default; - Range& operator=(Range&&) = default; - - KMM_HOST_DEVICE - constexpr Range() : begin(static_cast(0)), end(static_cast(0)) {} - - KMM_HOST_DEVICE - constexpr Range(T end) : begin(static_cast(0)), end(end) {} - - KMM_HOST_DEVICE - constexpr Range(T begin, T end) : begin(begin), end(end) {} - - template - KMM_HOST_DEVICE constexpr Range(const Range& that) { - if (!that.template is_convertible_to()) { - throw_overflow_exception(); - } - - *this = Range::from(that); - } - - template - KMM_HOST_DEVICE static Range from(const Range& range) { - return {static_cast(range.begin), static_cast(range.end)}; - } - - template - KMM_HOST_DEVICE constexpr bool is_convertible_to() const { - return is_convertible(begin) && is_convertible(end); - } - - /** - * Checks if the range is empty (i.e., `begin == end`) or invalid (i.e., `begin > end`). - */ - KMM_HOST_DEVICE - constexpr bool is_empty() const { - return this->begin >= this->end; - } - - /** - * Checks if the given index `index` is within this range. - */ - template - KMM_HOST_DEVICE constexpr bool contains(const U& index) const { - return !is_less(index, this->begin) && is_less(index, this->end); - } - - /** - * Checks if the given `that` range is fully contained within this range. - */ - template - KMM_HOST_DEVICE constexpr bool contains(const Range& that) const { - return that.is_empty() || // - (!is_less(that.begin, this->begin) && !is_less(this->end, that.end)); - } - - /** - * Checks if the given range `that` overlaps this range. - */ - template - KMM_HOST_DEVICE constexpr bool overlaps(const Range& that) const { - return this->begin < this->end && that.begin < that.end && // - is_less(this->begin, that.end) && is_less(that.begin, this->end); - } - - /** - * Returns the range that lies in the intersection of `this` and `that`. - */ - KMM_HOST_DEVICE - constexpr Range intersection(const Range& that) const { - return { - this->begin > that.begin ? this->begin : that.begin, - this->end < that.end ? this->end : that.end, - }; - } - - /** - * Computes the size (or length) of the range. - */ - KMM_HOST_DEVICE - constexpr T size() const { - return this->begin <= this->end ? this->end - this->begin : static_cast(0); - } - - /** - * Returns the range `mid...end` and modifies the current range such it becomes `begin...mid`. - */ - KMM_HOST_DEVICE - constexpr Range split_tail(T mid) { - if (mid < this->begin) { - mid = this->begin; - } - - if (mid > this->end) { - mid = this->end; - } - - auto old_end = this->end; - this->end = mid; - return {mid, old_end}; - } - - /** - * Returns a new range that has been shifted by the given amount. - */ - KMM_HOST_DEVICE - constexpr Range shift_by(T shift) const { - return {this->begin + shift, this->end + shift}; - } - - T begin; - T end; -}; - -template -Range(const T&) -> Range; - -template -Range(const T&, const T&) -> Range; - -template -KMM_HOST_DEVICE bool operator==(const Range& lhs, const Range& rhs) { - return is_equal(lhs.begin, rhs.begin) && is_equal(lhs.end, rhs.end); -} - -template -KMM_HOST_DEVICE bool operator!=(const Range& lhs, const Range& rhs) { - return !(lhs == rhs); -} -} // namespace kmm - -#if !KMM_IS_RTC - #include - #include - - #include "fmt/ostream.h" - - #include "kmm/utils/hash_utils.hpp" - -namespace kmm { - -template -std::ostream& operator<<(std::ostream& stream, const Range& p) { - return stream << p.begin << "..." << p.end; -} - -} // namespace kmm - -template -struct fmt::formatter>: fmt::ostream_formatter {}; - -template -struct std::hash> { - size_t operator()(const kmm::Range& p) const { - size_t result = 0; - kmm::hash_combine(result, p.begin); - kmm::hash_combine(result, p.end); - return result; - } -}; -#endif \ No newline at end of file diff --git a/include/kmm/utils/refcnt_ptr.hpp b/include/kmm/utils/refcnt_ptr.hpp new file mode 100644 index 00000000..8316e623 --- /dev/null +++ b/include/kmm/utils/refcnt_ptr.hpp @@ -0,0 +1,213 @@ +#pragma once + +#include +#include +#include +#include +#include + +namespace kmm { + +template +struct reference_count { + protected: + constexpr reference_count() noexcept = default; + + private: + template + friend struct refcnt_traits_impl; + + mutable std::atomic m_count {1}; +}; + +template +struct refcnt_traits_impl { + template>> + static void increment(const reference_count* that) noexcept { + that->m_count.fetch_add(1, std::memory_order_relaxed); + } + + template>> + static void decrement(const reference_count* that) noexcept { + if (that->m_count.fetch_sub(1, std::memory_order_acq_rel) == 1) { + delete static_cast(that); + } + } + + template>> + static size_t load_count(const reference_count* that) noexcept { + return that->m_count.load(std::memory_order_acquire); + } +}; + +template +struct refcnt_traits final: refcnt_traits_impl {}; + +/// Declares (without defining) an explicit specialization of `refcnt_traits`, so that +/// `refcnt_ptr` can be used in a header where `T` is still an incomplete type. Pair with +/// `KMM_REFCNT_TRAITS_IMPL(T)` in a source file where `T` is complete. +#define KMM_REFCNT_TRAITS_FWD(T) \ + template<> \ + struct refcnt_traits final { \ + static void increment(const T* that) noexcept; \ + static void decrement(const T* that) noexcept; \ + static size_t load_count(const T* that) noexcept; \ + }; + +#define KMM_REFCNT_TRAITS_IMPL(T) \ + void ::kmm::refcnt_traits::increment(const T* that) noexcept { \ + ::kmm::refcnt_traits_impl::increment(that); \ + } \ + void ::kmm::refcnt_traits::decrement(const T* that) noexcept { \ + ::kmm::refcnt_traits_impl::decrement(that); \ + } \ + size_t kmm::refcnt_traits::load_count(const T* that) noexcept { \ + return ::kmm::refcnt_traits_impl::load_count(that); \ + } + +template> +class refcnt_ptr { + using pointer = std::add_pointer_t; + using element_type = T; + using traits_type = Traits; + + public: + refcnt_ptr() noexcept = default; + + explicit refcnt_ptr(pointer p, bool increment_ref) noexcept : m_ptr(p) { + if (increment_ref && *this) { + traits_type::increment(this->get()); + } + } + + refcnt_ptr(std::unique_ptr&& other) noexcept : refcnt_ptr(other.release(), false) {} + + refcnt_ptr(const refcnt_ptr& other) noexcept : m_ptr(other.m_ptr) { + if (*this) { + traits_type::increment(this->get()); + } + } + + refcnt_ptr(refcnt_ptr&& other) noexcept : m_ptr(other.release()) {} + + template>> + refcnt_ptr(const refcnt_ptr& other) noexcept : m_ptr(other.get()) { + if (other) { + UTraits::increment(other.get()); + } + } + + template>> + refcnt_ptr(refcnt_ptr&& other) noexcept : m_ptr(other.release()) {} + + refcnt_ptr(std::nullptr_t) noexcept : refcnt_ptr() {} + + ~refcnt_ptr() { + if (*this) { + traits_type::decrement(this->get()); + } + } + + refcnt_ptr& operator=(const refcnt_ptr& other) noexcept { + if (&other != this) { + if (other) { + traits_type::increment(other.get()); + } + if (*this) { + traits_type::decrement(this->get()); + } + this->m_ptr = other.get(); + } + return *this; + } + + refcnt_ptr& operator=(refcnt_ptr&& other) noexcept { + if (&other != this) { + this->swap(other); + } + return *this; + } + + refcnt_ptr& operator=(std::nullptr_t) noexcept { + this->reset(); + return *this; + } + + element_type& operator*() const noexcept { + return *(this->get()); + } + + pointer operator->() const noexcept { + return this->get(); + } + + pointer get() const noexcept { + return this->m_ptr; + } + + explicit operator bool() const noexcept { + return this->get() != nullptr; + } + + /// Returns true if this is the only `refcnt_ptr` referring to the pointee. Note that this + /// is a snapshot: without external synchronization, another thread may concurrently copy or + /// drop a `refcnt_ptr` to the same object, invalidating the result immediately after it is read. + bool unique() const noexcept { + return *this && traits_type::load_count(this->get()) == 1; + } + + pointer release() noexcept { + auto* ptr = this->get(); + this->m_ptr = pointer {}; + return ptr; + } + + void reset(pointer p = pointer {}) noexcept { + *this = refcnt_ptr(p, false); + } + + void swap(refcnt_ptr& other) noexcept { + using std::swap; + swap(m_ptr, other.m_ptr); + } + + private: + pointer m_ptr = nullptr; +}; + +template +bool operator==(const refcnt_ptr& x, const refcnt_ptr& y) noexcept { + return x.get() == y.get(); +} + +template +bool operator!=(const refcnt_ptr& x, const refcnt_ptr& y) noexcept { + return !(x == y); +} + +template +bool operator==(const refcnt_ptr& x, std::nullptr_t) noexcept { + return x.get() == nullptr; +} + +template +bool operator!=(const refcnt_ptr& x, std::nullptr_t) noexcept { + return x.get() != nullptr; +} + +template +bool operator==(std::nullptr_t, const refcnt_ptr& x) noexcept { + return x.get() == nullptr; +} + +template +bool operator!=(std::nullptr_t, const refcnt_ptr& x) noexcept { + return x.get() != nullptr; +} + +template +refcnt_ptr make_refcnt(Args&&... args) { + return refcnt_ptr(new T(std::forward(args)...), false); +} + +} // namespace kmm diff --git a/include/kmm/utils/scope_exit.hpp b/include/kmm/utils/scope_exit.hpp new file mode 100644 index 00000000..933afb0a --- /dev/null +++ b/include/kmm/utils/scope_exit.hpp @@ -0,0 +1,57 @@ +#pragma once + +#include +#include + +#include "kmm/core/macros.hpp" + +namespace kmm { + +/// \addtogroup utility +/// @{ + +/** + * Invokes a callable when it goes out of scope, unless it has been released. + * + * Modeled after `std::experimental::scope_exit` from the Library Fundamentals TS. + * Useful for running cleanup code on every exit path from a scope, including + * early returns and exceptions. + * + * @code + * auto guard = scope_exit([&] { release(handle); }); + * // ... + * guard.release(); // optional: skip the cleanup + * @endcode + */ +template +class scope_exit { + KMM_NOT_COPYABLE_OR_MOVABLE(scope_exit) + + public: + explicit scope_exit(Fn fn) noexcept(std::is_nothrow_move_constructible_v) : + m_fn(std::move(fn)) {} + + ~scope_exit() { + if (m_active) { + m_fn(); + } + } + + /** + * Cancel the pending invocation, so the callable is not run on destruction. + */ + void release() noexcept { + m_active = false; + } + + private: + Fn m_fn; + bool m_active = true; +}; + +template +scope_exit(Fn) -> scope_exit; + +/// @} + +} // namespace kmm diff --git a/include/kmm/utils/small_vector.hpp b/include/kmm/utils/small_vector.hpp index a832be7f..b8e56876 100644 --- a/include/kmm/utils/small_vector.hpp +++ b/include/kmm/utils/small_vector.hpp @@ -4,15 +4,28 @@ #include #include #include +#include #include "fmt/ostream.h" +#include "kmm/core/panic.hpp" + namespace kmm { +/// \addtogroup utility +/// @{ + [[noreturn]] __attribute__((noinline)) void throw_small_vector_out_of_capacity(); template struct small_vector { + static_assert( + std::is_nothrow_default_constructible_v && std::is_nothrow_move_constructible_v + && std::is_nothrow_move_assignable_v, + "small_vector requires T to be nothrow default constructible, nothrow move constructible, " + "and nothrow move assignable" + ); + using capacity_type = uint32_t; small_vector() = default; @@ -51,6 +64,11 @@ struct small_vector { } small_vector& operator=(small_vector&& that) noexcept { + swap(that); + return *this; + } + + void swap(small_vector& that) noexcept { std::swap(this->m_inline_data, that.m_inline_data); std::swap(this->m_size, that.m_size); std::swap(this->m_capacity, that.m_capacity); @@ -63,11 +81,13 @@ struct small_vector { if (!that.is_heap_allocated()) { that.m_data = that.m_inline_data; } + } - return *this; + friend void swap(small_vector& a, small_vector& b) noexcept { + a.swap(b); } - ~small_vector() { + ~small_vector() noexcept { if (is_heap_allocated()) { delete[] m_data; } @@ -97,19 +117,20 @@ struct small_vector { return m_data; } - bool try_grow_capacity(size_t k = 1) { + bool try_grow_capacity(size_t k = 1) noexcept { capacity_type new_capacity = m_capacity; - do { + + if (new_capacity < 16) { + new_capacity = 16; + } + + while (new_capacity - m_size < k) { if (new_capacity > std::numeric_limits::max() - new_capacity) { // Failed to grow capacity return false; } new_capacity += new_capacity; - } while (new_capacity - m_size < k); - - if (new_capacity < 16) { - new_capacity = 16; } auto new_data = std::unique_ptr(new (std::nothrow_t {}) T[new_capacity] {}); @@ -124,6 +145,12 @@ struct small_vector { if (is_heap_allocated()) { delete[] m_data; + } else { + // The inline slots are not deallocated like a heap array would be, so drop + // their moved-from contents now instead of letting them linger unused. + for (size_t i = 0; i < m_size; i++) { + m_inline_data[i] = T {}; + } } m_capacity = new_capacity; @@ -131,7 +158,7 @@ struct small_vector { return true; } - bool try_push_back(T item) { + bool try_push_back(T item) noexcept { if (m_capacity <= m_size && !try_grow_capacity(1)) { return false; } @@ -151,14 +178,15 @@ struct small_vector { void insert_all(It begin, It end) { size_t n = static_cast(end - begin); - if (m_capacity - m_size < n) { + if (KMM_UNLIKELY(m_capacity - m_size < n)) { + // FIX: what if begin,end points to this small_vector, this will invalidate pointers if (!try_grow_capacity(n)) { throw_small_vector_out_of_capacity(); } } for (size_t i = 0; i < n; i++) { - m_data[m_size + i] = begin[i]; + m_data[m_size + i] = *(begin + i); } // This is safe since `n <= m_capacity - m_size` @@ -180,6 +208,11 @@ struct small_vector { } void resize(size_t n) { + if (n < m_size) { + truncate(n); + return; + } + if (m_capacity < n) { if (!try_grow_capacity(n - m_size)) { throw_small_vector_out_of_capacity(); @@ -190,17 +223,22 @@ struct small_vector { m_size = static_cast(n); } - void truncate(size_t new_size) { - if (new_size > m_size) { + void truncate(size_t new_size) noexcept { + if (new_size >= m_size) { return; } - // Safe since `new_size < m_size` + // Drop the trailing elements so any resource they hold is released now, + // rather than lingering until the slot is reused or the vector is destroyed. + for (size_t i = new_size; i < m_size; i++) { + m_data[i] = T {}; + } + m_size = static_cast(new_size); } void clear() noexcept { - m_size = 0; + truncate(0); } T& operator[](size_t i) noexcept { @@ -231,7 +269,7 @@ struct small_vector { T* m_data = m_inline_data; capacity_type m_size = 0; capacity_type m_capacity = InlineSize; - T m_inline_data[InlineSize]; + T m_inline_data[InlineSize] {}; }; using byte_buffer = small_vector; @@ -249,6 +287,9 @@ std::ostream& operator<<(std::ostream& stream, const small_vector& p) { return stream << "}"; } + +/// @} + } // namespace kmm template diff --git a/src/api/array.cpp b/src/api/array.cpp deleted file mode 100644 index 8284e5e7..00000000 --- a/src/api/array.cpp +++ /dev/null @@ -1,3 +0,0 @@ -#include "kmm/api/array.hpp" - -namespace kmm {} \ No newline at end of file diff --git a/src/api/array_instance.cpp b/src/api/array_instance.cpp deleted file mode 100644 index d30e0a1f..00000000 --- a/src/api/array_instance.cpp +++ /dev/null @@ -1,76 +0,0 @@ -#include "kmm/api/array_instance.hpp" -#include "kmm/runtime/runtime.hpp" - -namespace kmm { - -template -ArrayInstance::ArrayInstance( - TaskGraph& stage, - Runtime& rt, - Distribution dist, - DataType dtype -) : - ArrayDescriptor(stage, dist, dtype), - m_rt(rt.shared_from_this()) {} - -template -std::shared_ptr> ArrayInstance::create( - Runtime& rt, - Distribution dist, - DataType dtype -) { - std::shared_ptr> result; - - rt.schedule([&](auto& stage) { - result = std::shared_ptr>( - new ArrayInstance(stage, rt, std::move(dist), dtype) - ); - }); - - return result; -} - -template -ArrayInstance::~ArrayInstance() { - m_rt->schedule([&](TaskGraph& stage) { // - this->destroy(stage); - }); -} - -template -void ArrayInstance::copy_bytes_into(void* dst_data) { - EventId event_id = m_rt->schedule([&](TaskGraph& stage) { // - return this->copy_bytes_into_buffer(stage, dst_data); - }); - - m_rt->query_event(event_id, std::chrono::system_clock::time_point::max()); -} - -template -void ArrayInstance::copy_bytes_from(const void* src_data) { - EventId event_id = m_rt->schedule([&](TaskGraph& stage) { // - return this->copy_bytes_from_buffer(stage, src_data); - }); - - m_rt->query_event(event_id, std::chrono::system_clock::time_point::max()); -} - -template -void ArrayInstance::synchronize() const { - EventId event_id = m_rt->schedule([&](TaskGraph& stage) { // - return this->join_events(stage); - }); - - m_rt->query_event(event_id, std::chrono::system_clock::time_point::max()); -} - -[[noreturn]] void throw_uninitialized_array_exception() { - throw std::runtime_error( - "attempted to access an uninitialized array, " - "no associated instance was found" - ); -} - -KMM_INSTANTIATE_ARRAY_IMPL(ArrayInstance) - -} // namespace kmm \ No newline at end of file diff --git a/src/api/buffer.cpp b/src/api/buffer.cpp new file mode 100644 index 00000000..12328fad --- /dev/null +++ b/src/api/buffer.cpp @@ -0,0 +1,209 @@ +#include + +#include "spdlog/spdlog.h" + +#include "kmm/api/buffer.hpp" +#include "kmm/api/context.hpp" +#include "kmm/runtime/resource.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +struct Buffer::Impl: reference_count { + Impl(Runtime rt, BufferId id, BufferLayout layout) : + runtime(std::move(rt)), + id(id), + layout(layout) {} + + ~Impl() { + runtime.release_buffer(id); + } + + Runtime runtime; + BufferId id; + BufferLayout layout; +}; + +Buffer Buffer::create( + Runtime runtime, + BufferLayout layout, + std::string name, + FillValue fill_value, + std::optional home, + std::optional kind +) { + auto id = runtime.create_buffer(layout, std::move(name), fill_value, home, kind); + Buffer buffer; + buffer.m_impl = make_refcnt(runtime, id, layout); + return buffer; +} + +Buffer Buffer::adopt( + Runtime runtime, + BufferLayout layout, + void* ptr, + MemoryId memory_id, + std::string name +) { + auto id = runtime.adopt_buffer(layout, std::move(name), ptr, memory_id); + Buffer buffer; + buffer.m_impl = make_refcnt(runtime, id, layout); + return buffer; +} + +static void assert_impl(const refcnt_ptr& impl) { + if (impl == nullptr) { + throw std::runtime_error( + "cannot access buffer as it has not yet been registered with a runtime system" + ); + } +} + +BufferId Buffer::id() const { + assert_impl(m_impl); + return m_impl->id; +} + +Runtime Buffer::runtime() const { + assert_impl(m_impl); + return m_impl->runtime; +} + +BufferLayout Buffer::layout() const { + assert_impl(m_impl); + return m_impl->layout; +} + +std::optional Buffer::home() const { + assert_impl(m_impl); + return m_impl->runtime.buffer_home(m_impl->id); +} + +void Buffer::prefetch(MemoryId memory_id, bool invalidate_others) const { + if (m_impl) { + m_impl->runtime.prefetch_buffer(m_impl->id, memory_id, invalidate_others); + } +} + +void Buffer::poison(std::exception_ptr reason) const { + assert_impl(m_impl); + m_impl->runtime.poison_buffer(m_impl->id, std::move(reason)); +} + +void Buffer::invalidate() const { + assert_impl(m_impl); + m_impl->runtime.invalidate_buffer(id()); +} + +static void copy_device_to_host( + const void* src_addr, + void* dst_addr, + const CopyDescription& simplified +) { + KMM_ASSERT(simplified.num_dims == 0); + + auto src = static_cast(src_addr) + simplified.src_offset; + auto dst = static_cast(dst_addr) + simplified.dst_offset; + + KMM_GPU_CHECK(g_memcpy_d_to_h(dst, (g_device_ptr_t)src, simplified.element_size)); +} + +static void copy_host_to_device( + const void* src_addr, + void* dst_addr, + const CopyDescription& simplified +) { + KMM_ASSERT(simplified.num_dims == 0); + + auto src = static_cast(src_addr) + simplified.src_offset; + auto dst = static_cast(dst_addr) + simplified.dst_offset; + + KMM_GPU_CHECK(g_memcpy_h_to_d((g_device_ptr_t)dst, src, simplified.element_size)); +} + +static CopyDescription contiguous_copy_description( + size_t nbytes, + size_t src_offset, + size_t dst_offset +) { + CopyDescription description; + description.src_offset = checked_cast(src_offset); + description.dst_offset = checked_cast(dst_offset); + description.add_dimension(checked_cast(nbytes), 1, 1); + return description; +} + +void Buffer::copy_to(void* dest, size_t nbytes, size_t offset, MemoryId memory_id) const { + copy_to(dest, contiguous_copy_description(nbytes, offset, 0), memory_id); +} + +void Buffer::copy_to(void* dest, CopyDescription description, MemoryId memory_id) const { + assert_impl(m_impl); + auto runtime = m_impl->runtime; + + auto simplified = description.simplify(); + if (!memory_id.is_host() && simplified.num_dims != 0) { + spdlog::warn("copy could not be reduced to a 1D copy, falling back to a copy via host"); + memory_id = MemoryId::host(); + } + + ResourceRequest requests; + requests.add(memory_id, id(), AccessMode::Read); + auto grant = runtime.submit(std::move(requests)); + + auto accessor = grant.accessor(0); + KMM_ASSERT(accessor.memory_id == memory_id); + + auto src_range = description.src_range(); + KMM_ASSERT( + src_range.start >= 0 && static_cast(src_range.stop) <= accessor.size_in_bytes + ); + + if (memory_id.is_host()) { + memops::copy(accessor.address, dest, description); + } else { + copy_device_to_host(accessor.address, dest, simplified); + } + + runtime.release(grant); +} + +void Buffer::copy_from(const void* dest, size_t nbytes, size_t offset, MemoryId memory_id) const { + copy_from(dest, contiguous_copy_description(nbytes, 0, offset), memory_id); +} + +void Buffer::copy_from(const void* src, CopyDescription description, MemoryId memory_id) const { + assert_impl(m_impl); + auto runtime = m_impl->runtime; + + auto simplified = description.simplify(); + if (!memory_id.is_host() && simplified.num_dims != 0) { + spdlog::warn("copy could not be reduced to a 1D copy, falling back to a copy via host"); + memory_id = MemoryId::host(); + } + + ResourceRequest requests; + requests.add(memory_id, id(), AccessMode::ReadWrite); + auto grant = runtime.submit(std::move(requests)); + + auto accessor = grant.accessor(0); + KMM_ASSERT(accessor.memory_id == memory_id); + KMM_ASSERT(accessor.is_writable); + + auto dst_range = description.dst_range(); + KMM_ASSERT( + dst_range.start >= 0 && static_cast(dst_range.stop) <= accessor.size_in_bytes + ); + + if (memory_id.is_host()) { + memops::copy(src, accessor.address, description); + } else { + copy_host_to_device(src, accessor.address, simplified); + } + + runtime.release(grant); +} + +KMM_REFCNT_TRAITS_IMPL(Buffer::Impl) + +} // namespace kmm diff --git a/src/api/context.cpp b/src/api/context.cpp new file mode 100644 index 00000000..e69de29b diff --git a/src/api/device.cu b/src/api/device.cu new file mode 100644 index 00000000..92223f44 --- /dev/null +++ b/src/api/device.cu @@ -0,0 +1,22 @@ +#include "kmm/api/device.hpp" + +namespace kmm { + +Device::Device(Runtime runtime, DeviceId device_id, MemoryTransaction transaction) : + Context(runtime, std::move(transaction)), + m_device_id(device_id), + m_stream(nullptr) { + auto* context = runtime.system_info().device(device_id).context(); + m_stream = std::make_shared(context); +} + +DeviceStream Device::stream() const noexcept { + auto registry = Runtime(runtime()).event_registry(); + return {registry, registry.lookup_or_register_stream(*m_stream)}; +} + +Device Context::gpu(DeviceId device_id) { + return Device(m_runtime, device_id, m_transaction); +} + +} // namespace kmm diff --git a/src/api/host.cpp b/src/api/host.cpp new file mode 100644 index 00000000..7fea7d45 --- /dev/null +++ b/src/api/host.cpp @@ -0,0 +1,9 @@ +#include "kmm/api/host.hpp" + +namespace kmm { + +Host Context::host() { + return Host(m_runtime, m_transaction); +} + +} // namespace kmm diff --git a/src/api/mapper.cpp b/src/api/mapper.cpp deleted file mode 100644 index ed975107..00000000 --- a/src/api/mapper.cpp +++ /dev/null @@ -1,156 +0,0 @@ -#include -#include - -#include "kmm/api/mapper.hpp" -#include "kmm/utils/checked_math.hpp" - -namespace kmm { - -IndexMap::IndexMap(Axis variable, int64_t scale, int64_t offset, int64_t length, int64_t divisor) : - m_axis(variable), - m_scale(scale), - m_offset(offset), - m_length(length), - m_divisor(divisor) { - m_length = std::max(m_length, 0); - - if (m_divisor < 0) { - m_scale = -m_scale; - m_offset = -checked_add(m_offset, m_length - 1); - m_divisor = -m_divisor; - } - - if (m_scale != 1 && m_divisor != 1) { - auto common = std::gcd(std::gcd(m_scale, m_offset), std::gcd(m_length, m_divisor)); - - if (common != 1) { - m_scale /= common; - m_offset /= common; - m_length /= common; - m_divisor /= common; - } - } -} - -IndexMap IndexMap::range(IndexMap begin, IndexMap end) { - if (begin.m_scale != end.m_scale || begin.m_divisor != end.m_divisor) { - throw std::runtime_error(fmt::format( - "`range` requires two expressions having the same scaling factor, given: `{}` and `{}`", - begin, - end - )); - } - - if (begin.m_axis.get() != end.m_axis.get() && begin.m_scale != 0) { - throw std::runtime_error(fmt::format( - "`range` requires two expression operating on the same axis, given: `{}` and `{}`", - begin, - end - )); - } - - return { - begin.m_axis, - begin.m_scale, - begin.m_offset, - (end.m_offset - begin.m_offset) + end.m_length, - begin.m_divisor - }; -} - -IndexMap IndexMap::offset_by(int64_t offset) const { - auto new_offset = checked_add(m_offset, checked_mul(m_divisor, offset)); - return {m_axis, m_scale, new_offset, m_length, m_divisor}; -} - -IndexMap IndexMap::scale_by(int64_t factor) const { - if (factor < 0) { - return negate().scale_by(-factor); - } - - return { - m_axis, - checked_mul(m_scale, factor), - checked_mul(m_offset, factor), - checked_mul(m_length - 1, factor) + 1, - m_divisor - }; -} - -IndexMap IndexMap::divide_by(int64_t divisor) const { - if (divisor < 0) { - return negate().divide_by(-divisor); - } - - return {m_axis, m_scale, m_offset, m_length, checked_mul(m_divisor, divisor)}; -} - -IndexMap IndexMap::negate() const { - return {m_axis, -m_scale, -checked_add(m_offset, m_length - 1), m_length, m_divisor}; -} - -Bounds<1> IndexMap::apply(DomainChunk chunk) const { - int64_t a0 = chunk.offset.get_or_default(m_axis.get()); - int64_t an = chunk.size.get_or_default(m_axis.get()); - - int64_t b0; - int64_t bn; - - if (m_scale == 0) { - b0 = m_offset; - bn = m_length; - } else if (m_length <= 0 || an <= 0) { - b0 = 0; - bn = 0; - } else if (m_scale > 0) { - b0 = m_scale * a0 + m_offset; - bn = m_scale * (an - 1) + m_length; - } else { - b0 = m_scale * (a0 + an - 1) + m_offset; - bn = -m_scale * (an - 1) + m_length; - } - - if (m_divisor > 1) { - int64_t remainder = b0 % m_divisor; - b0 = b0 / m_divisor; - bn = (bn + remainder + m_divisor - 1) / m_divisor; - } - - return static_cast>(Range {b0, b0 + bn}); -} - -static void write_mapping(std::ostream& f, Axis v, int64_t scale, int64_t offset, int64_t divisor) { - static constexpr const char* variables[] = {"x", "y", "z", "w"}; - const char* var = v.get() < 4 ? variables[v.get()] : "?"; - - if (scale != 1) { - if (offset != 0) { - f << "(" << scale << "*" << var << " + " << offset << ")"; - } else { - f << scale << "*" << var; - } - } else { - if (offset != 0) { - f << "(" << var << " + " << offset << ")"; - } else { - f << var; - } - } - - if (divisor != 1) { - f << "/" << divisor; - } -} - -std::ostream& operator<<(std::ostream& f, const IndexMap& that) { - write_mapping(f, that.m_axis, that.m_scale, that.m_offset, that.m_divisor); - - if (that.m_length != 1) { - f << "..."; - write_mapping(f, that.m_axis, that.m_scale, that.m_offset + that.m_length, that.m_divisor); - } - - return f; -} - -} // namespace kmm \ No newline at end of file diff --git a/src/api/runtime_handle.cpp b/src/api/runtime_handle.cpp deleted file mode 100644 index 5450b426..00000000 --- a/src/api/runtime_handle.cpp +++ /dev/null @@ -1,89 +0,0 @@ -#include "kmm/api/runtime_handle.hpp" -#include "kmm/core/resource.hpp" -#include "kmm/runtime/runtime.hpp" - -namespace kmm { - -struct RuntimeHandle::Impl { - std::shared_ptr worker; - SystemInfo info; - - Impl(std::shared_ptr worker, SystemInfo info) : worker(worker), info(info) {} -}; - -RuntimeHandle::RuntimeHandle(std::shared_ptr impl) { - KMM_ASSERT(impl != nullptr && impl->worker != nullptr); - m_data = std::move(impl); -} - -RuntimeHandle::RuntimeHandle(std::shared_ptr rt) : - RuntimeHandle(std::make_shared(rt, rt->system_info())) {} - -RuntimeHandle::RuntimeHandle(Runtime& rt) : RuntimeHandle(rt.shared_from_this()) {} - -MemoryId RuntimeHandle::memory_affinity_for_address(const void* address) const { - if (auto device_opt = get_gpu_device_by_address(address)) { - const auto& device = worker().system_info().device_by_ordinal(*device_opt); - return device.memory_id(); - } else { - return MemoryId::host(); - } -} - -EventId RuntimeHandle::join(EventList events) const { - return worker().schedule([&](TaskGraph& g) { return g.join_events(std::move(events)); }); -} - -bool RuntimeHandle::is_done(EventId id) const { - return worker().query_event(id, std::chrono::system_clock::time_point::min()); -} - -void RuntimeHandle::wait(EventId id) const { - worker().query_event(id, std::chrono::system_clock::time_point::max()); -} - -bool RuntimeHandle::wait_until(EventId id, typename std::chrono::system_clock::time_point deadline) - const { - return worker().query_event(id, deadline); -} - -bool RuntimeHandle::wait_for(EventId id, typename std::chrono::system_clock::duration duration) - const { - return worker().query_event(id, std::chrono::system_clock::now() + duration); -} - -EventId RuntimeHandle::barrier() const { - return worker().schedule([&](TaskGraph& g) { // - return g.insert_barrier(); - }); -} - -void RuntimeHandle::synchronize() const { - wait(barrier()); -} - -RuntimeHandle RuntimeHandle::constrain_to(std::vector resources) const { - return std::make_shared(m_data->worker, SystemInfo(m_data->info, std::move(resources))); -} - -RuntimeHandle RuntimeHandle::constrain_to(DeviceId device) const { - return constrain_to(ResourceId(device)); -} - -RuntimeHandle RuntimeHandle::constrain_to(ResourceId resource) const { - return constrain_to(std::vector {resource}); -} - -const SystemInfo& RuntimeHandle::info() const { - return m_data->info; -} - -Runtime& RuntimeHandle::worker() const { - return *m_data->worker; -} - -RuntimeHandle make_runtime(const RuntimeConfig& config) { - return make_worker(config); -} - -} // namespace kmm diff --git a/src/backends/cuda.cpp b/src/backends/cuda.cpp deleted file mode 100644 index fe2feb16..00000000 --- a/src/backends/cuda.cpp +++ /dev/null @@ -1,26 +0,0 @@ -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -g_result_t g_memcpy_peer_async( - g_device_ptr_t dstDevicePtr, - g_context_t dstContext, - g_device_t dstDevice, - g_device_ptr_t srcDevicePtr, - g_context_t srcContext, - g_device_t srcDevice, - size_t ByteCount, - g_stream_t hStream -) { - return cuMemcpyPeerAsync( - dstDevicePtr, - dstContext, - srcDevicePtr, - srcContext, - ByteCount, - hStream - ); -} - -} // namespace kmm diff --git a/src/backends/hip.cpp b/src/backends/hip.cpp deleted file mode 100644 index 9ddbc231..00000000 --- a/src/backends/hip.cpp +++ /dev/null @@ -1,45 +0,0 @@ -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -const char* blas_get_status_name(blas_status_t) { - return ""; -} - -g_result_t g_memcpy_async( - g_device_ptr_t dst, - g_device_ptr_t src, - size_t ByteCount, - g_stream_t hStream -) { - return hipMemcpyAsync(dst, src, ByteCount, hipMemcpyDefault, hStream); -} - -g_result_t g_memcpy_h_to_d_async( - g_device_ptr_t dstDevice, - const void* srcHost, - size_t ByteCount, - g_stream_t hStream -) { - return hipMemcpyHtoDAsync(dstDevice, const_cast(srcHost), ByteCount, hStream); -} - -g_result_t g_memcpy_h_to_d(g_device_ptr_t dstDevice, const void* srcHost, size_t ByteCount) { - return hipMemcpyHtoD(dstDevice, const_cast(srcHost), ByteCount); -} - -g_result_t g_memcpy_peer_async( - g_device_ptr_t dstDevicePtr, - g_context_t dstContext, - g_device_t dstDevice, - g_device_ptr_t srcDevicePtr, - g_context_t srcContext, - g_device_t srcDevice, - size_t ByteCount, - g_stream_t hStream -) { - return hipMemcpyPeerAsync(dstDevicePtr, dstDevice, srcDevicePtr, srcDevice, ByteCount, hStream); -} - -} // namespace kmm diff --git a/src/core/backends.cpp b/src/core/backends.cpp deleted file mode 100644 index a781ae84..00000000 --- a/src/core/backends.cpp +++ /dev/null @@ -1,8 +0,0 @@ - -#ifdef KMM_USE_CUDA - #include "../backends/cuda.cpp" -#elif KMM_USE_HIP - #include "../backends/hip.cpp" -#else - #include "../backends/cpu.cpp" -#endif \ No newline at end of file diff --git a/src/utils/checked_compare.cpp b/src/core/checked_compare.cpp similarity index 58% rename from src/utils/checked_compare.cpp rename to src/core/checked_compare.cpp index 777c9d9e..99c93dd2 100644 --- a/src/utils/checked_compare.cpp +++ b/src/core/checked_compare.cpp @@ -1,10 +1,10 @@ #include -#include "kmm/utils/checked_compare.hpp" +#include "kmm/core/checked_compare.hpp" namespace kmm { -void throw_overflow_exception() { +[[noreturn]] void throw_overflow_exception() { throw std::overflow_error("integer overflow occurred"); } diff --git a/src/core/data_type.cpp b/src/core/data_type.cpp deleted file mode 100644 index f3b67735..00000000 --- a/src/core/data_type.cpp +++ /dev/null @@ -1,139 +0,0 @@ -#include - -#include "fmt/format.h" - -#include "kmm/core/data_type.hpp" -#include "kmm/utils/panic.hpp" - -namespace kmm { - -const char* scalar_name(ScalarType kind) { - switch (kind) { - case ScalarType::Int8: - return "Int8"; - case ScalarType::Int16: - return "Int16"; - case ScalarType::Int32: - return "Int32"; - case ScalarType::Int64: - return "Int64"; - case ScalarType::Uint8: - return "Uint8"; - case ScalarType::Uint16: - return "Uint16"; - case ScalarType::Uint32: - return "Uint32"; - case ScalarType::Uint64: - return "Uint64"; - case ScalarType::Float16: - return "Float16"; - case ScalarType::Float32: - return "Float32"; - case ScalarType::Float64: - return "Float64"; - case ScalarType::BFloat16: - return "BFloat16"; - case ScalarType::Complex16: - return "Float64"; - case ScalarType::Complex32: - return "Complex32"; - case ScalarType::Complex64: - return "Complex64"; - case ScalarType::KeyAndInt64: - return "KeyAndInt64"; - case ScalarType::KeyAndFloat64: - return "KeyAndFloat64"; - default: - return "(unknown type)"; - } -} - -DataType DataType::of(ScalarType kind) { - switch (kind) { - case ScalarType::Int8: - return DataType::of(); - case ScalarType::Int16: - return DataType::of(); - case ScalarType::Int32: - return DataType::of(); - case ScalarType::Int64: - return DataType::of(); - case ScalarType::Uint8: - return DataType::of(); - case ScalarType::Uint16: - return DataType::of(); - case ScalarType::Uint32: - return DataType::of(); - case ScalarType::Uint64: - return DataType::of(); - // case ScalarType::Float16: - // return DataType::of(); - case ScalarType::Float32: - return DataType::of(); - case ScalarType::Float64: - return DataType::of(); - // case ScalarType::BFloat16: - // break; - // case ScalarType::Complex16: - // break; - case ScalarType::Complex32: - return DataType::of>(); - case ScalarType::Complex64: - return DataType::of>(); - case ScalarType::KeyAndInt64: - return DataType::of>(); - case ScalarType::KeyAndFloat64: - return DataType::of>(); - case ScalarType::Invalid: - return DataType(); - } - - throw std::runtime_error(fmt::format("unknown scalar type: {}", scalar_name(kind))); -} - -const std::type_info& DataType::type_info() const { - KMM_ASSERT(m_info != nullptr); - return m_info->type_id; -} - -ScalarType DataType::as_scalar() const { - return m_info != nullptr ? m_info->scalar_type : ScalarType::Invalid; -} - -size_t DataType::size_in_bytes() const { - KMM_ASSERT(m_info != nullptr); - return m_info->size_in_bytes; -} - -size_t DataType::alignment() const { - KMM_ASSERT(m_info != nullptr); - return m_info->alignment; -} - -const char* DataType::name() const { - if (m_info == nullptr) { - return "Invalid"; - } else if (const auto* name = m_info->name) { - return name; - } else { - return m_info->type_id.name(); - } -} - -const char* DataType::c_name() const { - if (m_info != nullptr && m_info->c_name != nullptr) { - return m_info->c_name; - } - - return name(); -} - -std::ostream& operator<<(std::ostream& f, ScalarType p) { - return f << scalar_name(p); -} - -std::ostream& operator<<(std::ostream& f, DataType p) { - return f << p.name(); -} - -} // namespace kmm \ No newline at end of file diff --git a/src/core/distribution.cpp b/src/core/distribution.cpp deleted file mode 100644 index 04de0160..00000000 --- a/src/core/distribution.cpp +++ /dev/null @@ -1,205 +0,0 @@ -#include "kmm/core/distribution.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -template -Bounds index2region( - size_t index, - std::array num_chunks, - Dim chunk_size, - Dim array_size -) { - Bounds result; - - for (size_t j = 0; is_less(j, N); j++) { - size_t i = N - 1 - j; - auto k = index % num_chunks[i]; - index /= num_chunks[i]; - - result[i].begin = int64_t(k) * chunk_size[i]; - result[i].end = std::min(result[i].begin + chunk_size[i], array_size[i]); - } - - return result; -} - -template -size_t region2index( - Bounds region, - Dim chunk_size, - Dim array_size, - std::array chunks_count -) { - size_t index = 0; - - for (size_t i = 0; is_less(i, N); i++) { - auto k = div_floor(region.begin(i), chunk_size[i]); - - if (!in_range(k, chunks_count[i])) { - throw std::out_of_range(fmt::format( - "invalid read pattern, the region {} exceeds the array dimensions {}", - region, - array_size - )); - } - - if (region.end(i) > (k + 1) * chunk_size[i]) { - throw std::out_of_range(fmt::format( - "invalid read pattern, the region {} does not align to the chunk size of {}", - region, - chunk_size - )); - } - - index = index * chunks_count[i] + static_cast(k); - } - - return index; -} - -template -Distribution::Distribution() : m_array_size(Dim::zero()), m_chunk_size(Dim::one()) { - for (size_t i = 0; is_less(i, N); i++) { - m_chunks_count[i] = 0; - } - - // If `N==0`, there is always one chunk. Just assign it to the host for now. - if constexpr (N == 0) { - m_memories.push_back(MemoryId::host()); - } -} - -template -Distribution::Distribution( - Dim array_size, - Dim chunk_size, - std::vector memories -) : - m_array_size(array_size), - m_chunk_size(chunk_size), - m_memories(std::move(memories)) { - size_t total_chunk_count = 1; - - for (size_t i = 0; is_less(i, N); i++) { - m_chunks_count[i] = checked_cast(div_ceil(array_size[i], chunk_size[i])); - total_chunk_count = checked_mul(total_chunk_count, m_chunks_count[i]); - } - - if (total_chunk_count != m_memories.size()) { - throw std::runtime_error(fmt::format( - "data distribution contains {} chunk(s), only {} memory location(s) provided", - total_chunk_count, - m_memories.size() - )); - } -} - -template -Distribution Distribution::from_chunks( - Dim array_size, - std::vector> chunks, - bool allow_duplicates -) { - size_t total_chunk_count = 1; - std::array chunks_count; - Dim chunk_size = Dim::zero(); - - for (const auto& chunk : chunks) { - if (chunk.offset == Point::zero()) { - chunk_size = chunk.size; - } - } - - if (chunk_size.is_empty()) { - throw std::runtime_error("chunk size cannot be empty"); - } - - for (size_t i = 0; is_less(i, N); i++) { - chunks_count[i] = checked_cast(div_ceil(array_size[i], chunk_size[i])); - total_chunk_count = checked_mul(total_chunk_count, chunks_count[i]); - } - - auto memories = std::vector(total_chunk_count, MemoryId::host()); - auto visited = std::vector(total_chunk_count, false); - - for (size_t index = 0; index < chunks.size(); index++) { - const auto chunk = chunks[index]; - size_t linear_index = 0; - Point expected_offset; - Dim expected_size; - - for (size_t i = 0; is_less(i, N); i++) { - auto k = div_floor(chunk.offset[i], chunk_size[i]); - - expected_offset[i] = k * chunk_size[i]; - expected_size[i] = std::min(chunk_size[i], array_size[i] - expected_offset[i]); - - linear_index = linear_index * chunks_count[i] + static_cast(k); - } - - if (chunk.offset != expected_offset || chunk.size != expected_size) { - throw std::runtime_error(fmt::format( - "invalid write access pattern, region {} is not aligned to the chunk size of {}", - Bounds::from_offset_size(chunk.offset, chunk.size), - chunk_size - )); - } - - if (visited[linear_index]) { - if (!allow_duplicates) { - throw std::runtime_error(fmt::format( - "invalid write access pattern, region {} is written to by more than one task", - Bounds::from_offset_size(expected_offset, expected_size) - )); - } - - if (memories[linear_index] != chunk.owner_id) { - memories[linear_index] = MemoryId::host(); - } - } else { - visited[linear_index] = true; - memories[linear_index] = chunk.owner_id; - } - } - - for (size_t i = 0; i < total_chunk_count; i++) { - if (!visited[i]) { - auto region = index2region(i, chunks_count, chunk_size, array_size); - - throw std::runtime_error( - fmt::format("invalid write access pattern, no task writes to region {}", region) - ); - } - } - - return {array_size, chunk_size, std::move(memories)}; -} - -template -size_t Distribution::region_to_chunk_index(Bounds region) const { - return region2index(region, m_chunk_size, m_array_size, m_chunks_count); -} - -template -ArrayChunk Distribution::chunk(size_t index) const { - if (index >= m_memories.size()) { - throw std::runtime_error(fmt::format( - "chunk {} is out of range, there are only {} chunk(s)", - index, - m_memories.size() - )); - } - - auto region = index2region(index, m_chunks_count, m_chunk_size, m_array_size); - - return ArrayChunk { - .owner_id = m_memories[index], - .offset = region.begin(), - .size = region.size() - }; -} - -KMM_INSTANTIATE_ARRAY_IMPL(Distribution) - -} // namespace kmm \ No newline at end of file diff --git a/src/core/domain.cpp b/src/core/domain.cpp deleted file mode 100644 index 59d75d84..00000000 --- a/src/core/domain.cpp +++ /dev/null @@ -1,66 +0,0 @@ -#include "spdlog/spdlog.h" - -#include "kmm/core/domain.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -Domain TileDomain::operator()(const SystemInfo& info, ExecutionSpace space) const { - std::vector devices; - - if (space == ExecutionSpace::Host) { - for (const auto& resource : info.resources()) { - devices.push_back(ResourceId::host(resource.device_affinity())); - } - - if (devices.empty()) { - devices.push_back(ResourceId::host()); - } - } else if (space == ExecutionSpace::Device) { - devices = info.resources(); - } - - if (devices.empty()) { - throw std::runtime_error("cannot partition work, no devices found"); - } - - std::vector chunks; - - if (m_domain_size.is_empty()) { - return {chunks}; - } - - if (m_tile_size.is_empty()) { - throw std::runtime_error(fmt::format("invalid chunk size: {}", m_tile_size)); - } - - std::array num_chunks; - - for (size_t i = 0; i < DOMAIN_DIMS; i++) { - num_chunks[i] = div_ceil(m_domain_size[i].size(), m_tile_size[i]); - } - - size_t owner_id = 0; - auto offset = DomainPoint {}; - auto size = DomainDim {}; - - for (int64_t z = 0; z < num_chunks[2]; z++) { - for (int64_t y = 0; y < num_chunks[1]; y++) { - for (int64_t x = 0; x < num_chunks[0]; x++) { - auto current = Point<3> {x, y, z}; - - for (size_t i = 0; i < DOMAIN_DIMS; i++) { - offset[i] = m_domain_size[i].begin + current[i] * m_tile_size[i]; - size[i] = std::min(m_tile_size[i], m_domain_size[i].end - offset[i]); - } - - chunks.push_back({devices[owner_id], offset, size}); - owner_id = (owner_id + 1) % devices.size(); - } - } - } - - return {std::move(chunks)}; -} - -} // namespace kmm diff --git a/src/core/identifiers.cpp b/src/core/identifiers.cpp deleted file mode 100644 index 4034490b..00000000 --- a/src/core/identifiers.cpp +++ /dev/null @@ -1,73 +0,0 @@ -#include -#include - -#include "kmm/core/identifiers.hpp" - -namespace kmm { - -std::ostream& operator<<(std::ostream& f, const NodeId& v) { - // Must cast to uin32_t, since uint8_t is formatted as a character (`char`) - return f << uint32_t(v.get()); -} - -std::ostream& operator<<(std::ostream& f, const DeviceId& v) { - // Must cast to uin32_t, since uint8_t is formatted as a character (`char`) - return f << uint32_t(v.get()); -} - -std::ostream& operator<<(std::ostream& f, const MemoryId& v) { - if (v.is_host()) { - return f << "RAM"; - } else { - return f << "GPU:" << v.as_device(); - } -} - -std::ostream& operator<<(std::ostream& f, const ResourceId& v) { - if (v.is_host()) { - return f << "CPU"; - } else { - return f << "GPU:" << v.as_device(); - } -} - -std::ostream& operator<<(std::ostream& f, const BufferId& v) { - return f << v.get(); -} - -std::ostream& operator<<(std::ostream& f, const EventId& v) { - return f << v.get(); -} - -std::ostream& operator<<(std::ostream& f, const EventList& v) { - if (v.size() == 0) { - return f << "[]"; - } - - if (v.size() == 1) { - return f << "[" << v[0] << "]"; - } - - std::vector events = {v.begin(), v.end()}; - std::sort(events.begin(), events.end()); - - auto it = std::unique(events.begin(), events.end()); - events.erase(it, events.end()); - - f << "["; - bool is_first = true; - - for (const auto& e : events) { - if (!is_first) { - f << ", "; - } - - is_first = false; - f << e; - } - - f << "]"; - return f; -} - -} // namespace kmm \ No newline at end of file diff --git a/src/utils/panic.cpp b/src/core/panic.cpp similarity index 63% rename from src/utils/panic.cpp rename to src/core/panic.cpp index 143642af..61e2ec4e 100644 --- a/src/utils/panic.cpp +++ b/src/core/panic.cpp @@ -1,5 +1,5 @@ -#include -#include +#include +#include // For POSIX stack trace (Linux, macOS, etc.) // Remove if not available on your platform @@ -9,28 +9,26 @@ namespace kmm { -void panic(const char* file, int line, const char* message) { +[[noreturn]] void panic(const char* file, int line, const char* message) { fprintf(stderr, "\nPANIC TRIGGERED\n"); - fprintf(stderr, " location: %s:%d\n", file, line); - fprintf(stderr, " message: %s\n", message); + fprintf(stderr, " location: %s:%d\n", file, line); + fprintf(stderr, " message: %s\n", message); + fprintf(stderr, " stack trace:\n"); #ifdef __GNUC__ // Attempt to capture and print a backtrace - void* callstack[128]; - int nframes = backtrace(callstack, 128); + void* callstack[256]; + int nframes = backtrace(callstack, 256); char** symbols = backtrace_symbols(callstack, nframes); - if (symbols != NULL) { - fprintf(stderr, " stack trace:\n"); + if (symbols != nullptr) { for (int i = 0; i < nframes; i++) { fprintf(stderr, " %s\n", symbols[i]); } } else { - fprintf(stderr, " stack trace:\n"); fprintf(stderr, " ??? \n"); } #else - fprintf(stderr, " stack trace:\n"); fprintf(stderr, " ??? \n"); #endif @@ -40,4 +38,4 @@ void panic(const char* file, int line, const char* message) { abort(); } -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/src/core/reduction.cpp b/src/core/reduction.cpp deleted file mode 100644 index 1ff2bf30..00000000 --- a/src/core/reduction.cpp +++ /dev/null @@ -1,28 +0,0 @@ -#include "kmm/core/reduction.hpp" - -namespace kmm { - -std::ostream& operator<<(std::ostream& f, Reduction p) { - switch (p) { - case Reduction::Sum: - return f << "Sum"; - case Reduction::Product: - return f << "Product"; - case Reduction::Min: - return f << "Min"; - case Reduction::Max: - return f << "Max"; - case Reduction::BitAnd: - return f << "BitAnd"; - case Reduction::BitOr: - return f << "BitOr"; - default: - return f << "(unknown operation)"; - } -} - -std::ostream& operator<<(std::ostream& f, ReductionOutput p) { - return f << "Reduction(" << p.operation << ", " << p.data_type << ")"; -} - -} // namespace kmm \ No newline at end of file diff --git a/src/core/resource.cpp b/src/core/resource.cpp deleted file mode 100644 index c4afd83c..00000000 --- a/src/core/resource.cpp +++ /dev/null @@ -1,72 +0,0 @@ -#include - -#include "fmt/format.h" -#include "spdlog/spdlog.h" - -#include "kmm/core/resource.hpp" -#include "kmm/memops/gpu_fill.hpp" -#include "kmm/utils/checked_math.hpp" - -namespace kmm { - -InvalidResourceException::InvalidResourceException( - const std::type_info& expected, - const std::type_info& gotten -) { - m_message = fmt::format( - "task expected an execution context of type {}, but was executed with type {}", - expected.name(), - gotten.name() - ); -} - -const char* InvalidResourceException::what() const noexcept { - return m_message.c_str(); -} - -DeviceResource::DeviceResource(DeviceInfo info, GPUContextHandle context, g_stream_t stream) : - DeviceInfo(info), - m_context(context), - m_stream(stream) { - GPUContextGuard guard {m_context}; - - KMM_GPU_CHECK(blas_create(&m_blas_handle)); - KMM_GPU_CHECK(blas_set_stream(m_blas_handle, m_stream)); -} - -DeviceResource::~DeviceResource() { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(blas_destroy(m_blas_handle)); -} - -void DeviceResource::synchronize() const { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(g_stream_synchronize(nullptr)); - KMM_GPU_CHECK(g_stream_synchronize(m_stream)); -} - -void DeviceResource::fill_bytes( - void* dest_buffer, - size_t nbytes, - const void* fill_pattern, - size_t fill_pattern_size -) const { - GPUContextGuard guard {m_context}; - execute_gpu_fill_async( - m_stream, - reinterpret_cast(dest_buffer), - FillDef(fill_pattern_size, nbytes / fill_pattern_size, fill_pattern) - ); -} - -void DeviceResource::copy_bytes(const void* source_buffer, void* dest_buffer, size_t nbytes) const { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(g_memcpy_async( - reinterpret_cast(dest_buffer), - reinterpret_cast(const_cast(source_buffer)), - nbytes, - m_stream - )); -} - -} // namespace kmm \ No newline at end of file diff --git a/src/core/system_info.cpp b/src/core/system_info.cpp deleted file mode 100644 index bbb0d78c..00000000 --- a/src/core/system_info.cpp +++ /dev/null @@ -1,147 +0,0 @@ -#include - -#include "fmt/format.h" - -#include "kmm/core/system_info.hpp" - -namespace kmm { - -DeviceInfo::DeviceInfo(DeviceId id, GPUContextHandle context, size_t num_concurrent_streams) : - m_id(id), - m_concurrent_stream(num_concurrent_streams) { - GPUContextGuard guard {context}; - - KMM_GPU_CHECK(g_ctx_get_device(&m_device_id)); - - char name[1024]; - KMM_GPU_CHECK(g_device_get_name(name, 1024, m_device_id)); - m_name = std::string(name); - - for (size_t i = 1; i < NUM_ATTRIBUTES; i++) { - auto attr = g_device_attribute_t(i); - KMM_GPU_CHECK(g_device_get_attribute(&m_attributes[i], attr, m_device_id)); - } - - size_t ignore_free_memory; - KMM_GPU_CHECK(g_mem_get_info(&ignore_free_memory, &m_memory_capacity)); -} - -dim3 DeviceInfo::max_block_dim() const { - return dim3( - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_X)), - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Y)), - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Z)) - ); -} - -dim3 DeviceInfo::max_grid_dim() const { - return dim3( - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_X)), - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Y)), - checked_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z)) - ); -} - -int DeviceInfo::compute_capability() const { - return (attribute(G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR) * 10) - + attribute(G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR); -} - -int DeviceInfo::max_threads_per_block() const { - return attribute(G_DEVICE_ATTRIBUTE_MAX_THREADS_PER_BLOCK); -} - -int DeviceInfo::attribute(g_device_attribute_t attrib) const { - if (attrib < NUM_ATTRIBUTES) { - return m_attributes[attrib]; - } - - throw std::runtime_error("unsupported attribute requested"); -} - -SystemInfo::SystemInfo(std::vector devices) : m_devices(devices) { - for (const auto& dev : devices) { - m_resources.push_back(ResourceId(dev.device_id())); - } -} - -SystemInfo::SystemInfo(SystemInfo base, std::vector subresources) : - m_devices(base.m_devices), - m_resources(std::move(subresources)) { - for (const auto& resource : m_resources) { - bool is_valid = false; - - for (auto base_resource : base.m_resources) { - is_valid |= base_resource.contains(resource); - } - - if (!is_valid) { - throw std::runtime_error(fmt::format("invalid resource: {}", resource)); - } - } -} - -size_t SystemInfo::num_devices() const { - return m_devices.size(); -} - -const DeviceInfo& SystemInfo::device(DeviceId id) const { - return m_devices.at(id.get()); -} - -const DeviceInfo& SystemInfo::device_by_ordinal(g_device_t ordinal) const { - for (const auto& device : m_devices) { - if (device.device_ordinal() == ordinal) { - return device; - } - } - - throw std::runtime_error(fmt::format("cannot find device with ordinal {}", ordinal)); -} - -std::vector SystemInfo::resources() const { - return m_resources; -} - -std::vector SystemInfo::memories() const { - std::vector result {MemoryId::host()}; - for (const auto& device : m_devices) { - result.push_back(device.memory_id()); - } - - return result; -} - -MemoryId SystemInfo::affinity_memory(DeviceId device_id) const { - return device(device_id).memory_id(); -} - -MemoryId SystemInfo::affinity_memory(ResourceId proc_id) const { - if (proc_id.is_device()) { - return affinity_memory(proc_id.as_device()); - } else { - return MemoryId::host(); - } -} - -ResourceId SystemInfo::affinity_processor(MemoryId memory_id) { - if (memory_id.is_device()) { - return memory_id.as_device(); - } else { - return ResourceId::host(); - } -} - -bool SystemInfo::is_memory_accessible(MemoryId memory_id, ResourceId proc_id) const { - if (!memory_id.is_host() && proc_id.is_device()) { - return affinity_memory(proc_id.as_device()) == memory_id; - } - - return memory_id.is_host(); -} - -bool SystemInfo::is_memory_accessible(MemoryId memory_id, DeviceId device_id) const { - return is_memory_accessible(memory_id, ResourceId(device_id)); -} - -} // namespace kmm diff --git a/src/core/task.cpp b/src/core/task.cpp deleted file mode 100644 index 975beeae..00000000 --- a/src/core/task.cpp +++ /dev/null @@ -1 +0,0 @@ -namespace kmm {} // namespace kmm \ No newline at end of file diff --git a/src/memops/gpu_copy.cpp b/src/memops/gpu_copy.cpp deleted file mode 100644 index 02e92923..00000000 --- a/src/memops/gpu_copy.cpp +++ /dev/null @@ -1,208 +0,0 @@ - -#include - -#include "fmt/format.h" - -#include "kmm/memops/types.hpp" -#include "kmm/utils/checked_math.hpp" -#include "kmm/utils/gpu_utils.hpp" - -namespace kmm { - -void throw_unsupported_dimension_exception(size_t dim) { - throw std::runtime_error(fmt::format( - "copy operation is {} dimensional, only 1D or 2D copy operations are supported", - dim + 1 - )); -} - -void execute_gpu_h2d_copy_impl( - std::optional stream, - const void* src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - copy_description.simplify(); - size_t dim = copy_description.effective_dimensionality(); - - g_device_ptr_t dst_ptr = gpu_deviceptr_offset(dst_buffer, copy_description.dst_offset); - const void* src_ptr = static_cast(src_buffer) + copy_description.src_offset; - - if (dim == 0) { - if (stream) { - KMM_GPU_CHECK(g_memcpy_h_to_d_async( // - dst_ptr, - src_ptr, - copy_description.element_size, - *stream - )); - } else { - KMM_GPU_CHECK(g_memcpy_h_to_d( // - dst_ptr, - src_ptr, - copy_description.element_size - )); - } - } else if (dim == 1) { - gpu_memcpy2d_t info; - ::memset(&info, 0, sizeof(gpu_memcpy2d_t)); - - info.srcMemoryType = g_memory_type_t::G_MEMORYTYPE_HOST; - info.srcHost = src_ptr; - info.srcPitch = checked_cast(copy_description.src_strides[0]); - info.dstMemoryType = g_memory_type_t::G_MEMORYTYPE_DEVICE; - info.dstDevice = dst_ptr; - info.dstPitch = checked_cast(copy_description.dst_strides[0]); - info.WidthInBytes = checked_cast(copy_description.element_size); - info.Height = checked_cast(copy_description.counts[0]); - - if (stream) { - KMM_GPU_CHECK(g_memcpy_2d_async(&info, *stream)); - } else { - KMM_GPU_CHECK(g_memcpy_2d(&info)); - } - } else { - throw_unsupported_dimension_exception(dim); - } -} - -void execute_gpu_d2h_copy_impl( - std::optional stream, - g_device_ptr_t src_buffer, - void* dst_buffer, - CopyDef copy_description -) { - copy_description.simplify(); - size_t dim = copy_description.effective_dimensionality(); - - void* dst_ptr = static_cast(dst_buffer) + copy_description.dst_offset; - g_device_ptr_t src_ptr = gpu_deviceptr_offset(src_buffer, copy_description.src_offset); - - if (dim == 0) { - if (stream) { - KMM_GPU_CHECK( - g_memcpy_d_to_h_async(dst_ptr, src_ptr, copy_description.element_size, *stream) - ); - } else { - KMM_GPU_CHECK(g_memcpy_d_to_h(dst_ptr, src_ptr, copy_description.element_size)); - } - } else if (dim == 1) { - gpu_memcpy2d_t info; - ::memset(&info, 0, sizeof(gpu_memcpy2d_t)); - - info.srcMemoryType = g_memory_type_t::G_MEMORYTYPE_DEVICE; - info.srcDevice = src_ptr; - info.srcPitch = checked_cast(copy_description.src_strides[0]); - info.dstMemoryType = g_memory_type_t::G_MEMORYTYPE_HOST; - info.dstHost = dst_ptr; - info.dstPitch = checked_cast(copy_description.dst_strides[0]); - info.WidthInBytes = checked_cast(copy_description.element_size); - info.Height = checked_cast(copy_description.counts[0]); - - if (stream) { - KMM_GPU_CHECK(g_memcpy_2d_async(&info, *stream)); - } else { - KMM_GPU_CHECK(g_memcpy_2d(&info)); - } - } else { - throw_unsupported_dimension_exception(dim); - } -} - -void execute_gpu_d2d_copy_impl( - std::optional stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - copy_description.simplify(); - size_t dim = copy_description.effective_dimensionality(); - - g_device_ptr_t dst_ptr = gpu_deviceptr_offset(dst_buffer, copy_description.dst_offset); - g_device_ptr_t src_ptr = gpu_deviceptr_offset(src_buffer, copy_description.src_offset); - - if (dim == 0) { - if (stream) { - KMM_GPU_CHECK(g_memcpy_d_to_d_async( // - dst_ptr, - src_ptr, - copy_description.element_size, - *stream - )); - } else { - KMM_GPU_CHECK(g_memcpy_d_to_d( // - dst_ptr, - src_ptr, - copy_description.element_size - )); - } - } else if (dim == 1) { - gpu_memcpy2d_t info; - ::memset(&info, 0, sizeof(gpu_memcpy2d_t)); - - info.srcMemoryType = g_memory_type_t::G_MEMORYTYPE_DEVICE; - info.srcDevice = src_ptr; - info.srcPitch = checked_cast(copy_description.src_strides[0]); - info.dstMemoryType = g_memory_type_t::G_MEMORYTYPE_DEVICE; - info.dstDevice = dst_ptr; - info.dstPitch = checked_cast(copy_description.dst_strides[0]); - info.WidthInBytes = checked_cast(copy_description.element_size); - info.Height = checked_cast(copy_description.counts[0]); - - if (stream) { - KMM_GPU_CHECK(g_memcpy_2d_async(&info, *stream)); - } else { - KMM_GPU_CHECK(g_memcpy_2d(&info)); - } - } else { - throw_unsupported_dimension_exception(dim); - } -} - -void execute_gpu_h2d_copy( - const void* src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - execute_gpu_h2d_copy_impl(std::nullopt, src_buffer, dst_buffer, copy_description); -} - -void execute_gpu_h2d_copy_async( - g_stream_t stream, - const void* src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - execute_gpu_h2d_copy_impl(stream, src_buffer, dst_buffer, copy_description); -} - -void execute_gpu_d2h_copy(g_device_ptr_t src_buffer, void* dst_buffer, CopyDef copy_description) { - execute_gpu_d2h_copy_impl(std::nullopt, src_buffer, dst_buffer, copy_description); -} - -void execute_gpu_d2h_copy_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - void* dst_buffer, - CopyDef copy_description -) { - execute_gpu_d2h_copy_impl(stream, src_buffer, dst_buffer, copy_description); -} - -void execute_gpu_d2d_copy( - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - execute_gpu_d2d_copy_impl(std::nullopt, src_buffer, dst_buffer, copy_description); -} - -void execute_gpu_d2d_copy_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - CopyDef copy_description -) { - execute_gpu_d2d_copy_impl(stream, src_buffer, dst_buffer, copy_description); -} -} // namespace kmm \ No newline at end of file diff --git a/src/memops/gpu_fill.cu b/src/memops/gpu_fill.cu deleted file mode 100644 index cae955a4..00000000 --- a/src/memops/gpu_fill.cu +++ /dev/null @@ -1,110 +0,0 @@ - -#include "spdlog/spdlog.h" - -#include "kmm/memops/gpu_fill.hpp" -#include "kmm/utils/gpu_utils.hpp" -#include "kmm/utils/integer_fun.hpp" -#include "kmm/utils/panic.hpp" - -namespace kmm { - -template -__global__ void fill_kernel(size_t nelements, T* dest_buffer, T fill_value) { - size_t i = blockIdx.x * size_t(block_size) + threadIdx.x; - - while (i < nelements) { - dest_buffer[i] = fill_value; - i += size_t(block_size) * gridDim.x; - } -} - -template -void submit_fill_kernel( - g_stream_t stream, - g_device_ptr_t dest_buffer, - size_t nelements, - const void* fill_pattern -) { - static constexpr uint32_t max_grid_size = 512; - static constexpr uint32_t block_size = 256; - - T fill_value; - ::memcpy(&fill_value, fill_pattern, sizeof(T)); - - uint32_t grid_size = nelements < max_grid_size * size_t(block_size) - ? static_cast(div_ceil(nelements, size_t(block_size))) - : max_grid_size; - - fill_kernel - <<>>(nelements, (T*)dest_buffer, fill_value); -} - -template -bool is_fill_pattern_repetitive(const void* fill_pattern, size_t fill_pattern_size) { - if (fill_pattern_size % N != 0) { - return false; - } - - for (size_t i = 1; i < fill_pattern_size / N; i++) { - for (size_t j = 0; j < N; j++) { - if (static_cast(fill_pattern)[i * N + j] - != static_cast(fill_pattern)[j]) { - return false; - } - } - } - - return true; -} - -void execute_gpu_fill_async(g_stream_t stream, g_device_ptr_t dest_buffer, const FillDef& fill) { - size_t element_size = fill.fill_value.size(); - size_t nbytes = fill.num_elements * element_size; - const void* fill_pattern = fill.fill_value.data(); - - if (nbytes == 0 || element_size == 0) { - return; - } - - if (is_fill_pattern_repetitive<1>(fill_pattern, element_size)) { - uint8_t pattern; - ::memcpy(&pattern, fill_pattern, sizeof(uint8_t)); - KMM_GPU_CHECK(g_memset_d8_async( // - g_device_ptr_t(dest_buffer), - pattern, - nbytes, - stream - )); - - } else if (is_fill_pattern_repetitive<2>(fill_pattern, element_size)) { - uint16_t pattern; - ::memcpy(&pattern, fill_pattern, sizeof(uint16_t)); - KMM_GPU_CHECK(g_memset_d16_async( // - g_device_ptr_t(dest_buffer), - pattern, - nbytes / sizeof(uint16_t), - stream - )); - - } else if (is_fill_pattern_repetitive<4>(fill_pattern, element_size)) { - uint32_t pattern; - ::memcpy(&pattern, fill_pattern, sizeof(uint32_t)); - KMM_GPU_CHECK(g_memset_d32_async( // - g_device_ptr_t(dest_buffer), - pattern, - nbytes / sizeof(uint32_t), - stream - )); - - } else if (is_fill_pattern_repetitive<8>(fill_pattern, element_size)) { - KMM_ASSERT((unsigned long long)(dest_buffer) % 8 == 0); // must be aligned? - submit_fill_kernel(stream, dest_buffer, nbytes / sizeof(uint64_t), fill_pattern); - } else { - throw GPUException(fmt::format( - "could not fill buffer, value is {} bits, but only 8, 16, 32 or 64 bit is supported", - element_size * 8 - )); - } -} - -} // namespace kmm diff --git a/src/memops/gpu_operators.cuh b/src/memops/gpu_operators.cuh deleted file mode 100644 index aca7f40a..00000000 --- a/src/memops/gpu_operators.cuh +++ /dev/null @@ -1,126 +0,0 @@ -#pragma once - -#include -#include - -#include "host_operators.hpp" - -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" - -namespace kmm { - -template -struct GPUAtomic; - -template -KMM_DEVICE void gpu_generic_atomicCAS(T* output, T input, F combine) { - static_assert(sizeof(T) == sizeof(M)); - static_assert(alignof(T) >= alignof(M)); - - M old_bits = *reinterpret_cast(output); - M assumed_bits; - M new_bits; - - do { - assumed_bits = old_bits; - - T old_value; - ::memcpy(&old_value, &old_bits, sizeof(T)); - - T new_value = combine(old_value, input); - ::memcpy(&new_bits, &new_value, sizeof(T)); - - if (assumed_bits == new_bits) { - break; - } - - old_bits = atomicCAS(reinterpret_cast(output), assumed_bits, new_bits); - } while (old_bits != assumed_bits); -} - -#ifndef KMM_USE_HIP -// TODO: add it back when HIP support will be present -template -struct GPUAtomic, std::enable_if_t> { - static KMM_DEVICE void atomic_combine(T* output, T input) { - gpu_generic_atomicCAS(output, input, ReductionOperator()); - } -}; -#endif - -template -struct GPUAtomic, std::enable_if_t> { - static KMM_DEVICE void atomic_combine(T* output, T input) { - gpu_generic_atomicCAS(output, input, ReductionOperator()); - } -}; - -template -struct GPUAtomic, std::enable_if_t> { - static KMM_DEVICE void atomic_combine(T* output, T input) { - gpu_generic_atomicCAS(output, input, ReductionOperator()); - } -}; - -#define KMM_GPU_ATOMIC_REDUCTION_IMPL(T, OP, EXPR) \ - template<> \ - struct GPUAtomic> { \ - static KMM_DEVICE void atomic_combine(T* output, T input) { \ - EXPR(output, input); \ - } \ - }; - -KMM_GPU_ATOMIC_REDUCTION_IMPL(int, Reduction::BitAnd, atomicAnd) -#ifndef KMM_USE_HIP -// TODO: add it back when HIP support will be present -KMM_GPU_ATOMIC_REDUCTION_IMPL(long long int, Reduction::BitAnd, atomicAnd) -#endif -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned int, Reduction::BitAnd, atomicAnd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned long long int, Reduction::BitAnd, atomicAnd) - -KMM_GPU_ATOMIC_REDUCTION_IMPL(int, Reduction::BitOr, atomicOr) -#ifndef KMM_USE_HIP -// TODO: add it back when HIP support will be present -KMM_GPU_ATOMIC_REDUCTION_IMPL(long long int, Reduction::BitOr, atomicOr) -#endif -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned int, Reduction::BitOr, atomicOr) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned long long int, Reduction::BitOr, atomicOr) - -KMM_GPU_ATOMIC_REDUCTION_IMPL(double, Reduction::Sum, atomicAdd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(float, Reduction::Sum, atomicAdd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(int, Reduction::Sum, atomicAdd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned int, Reduction::Sum, atomicAdd) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(long long int, ReductionOp::Sum, atomicAdd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned long long int, Reduction::Sum, atomicAdd) -#ifndef KMM_USE_HIP -// TODO: add it back when HIP support will be present -KMM_GPU_ATOMIC_REDUCTION_IMPL(half_type, Reduction::Sum, atomicAdd) -KMM_GPU_ATOMIC_REDUCTION_IMPL(bfloat16_type, Reduction::Sum, atomicAdd) -#endif - -//KMM_GPU_ATOMIC_REDUCTION_IMPL(double, ReductionOp::Min, atomicMin) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(float, ReductionOp::Min, atomicMin) -KMM_GPU_ATOMIC_REDUCTION_IMPL(int, Reduction::Min, atomicMin) -KMM_GPU_ATOMIC_REDUCTION_IMPL(long long int, Reduction::Min, atomicMin) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned int, Reduction::Min, atomicMin) -KMM_GPU_ATOMIC_REDUCTION_IMPL(long long unsigned int, Reduction::Min, atomicMin) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(__half, ReductionOp::Min, atomicMin) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(__nv_bfloat16, ReductionOp::Min, atomicMin) - -//KMM_GPU_ATOMIC_REDUCTION_IMPL(double, ReductionOp::Max, atomicMax) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(float, ReductionOp::Max, atomicMax) -KMM_GPU_ATOMIC_REDUCTION_IMPL(int, Reduction::Max, atomicMax) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned int, Reduction::Max, atomicMax) -KMM_GPU_ATOMIC_REDUCTION_IMPL(long long int, Reduction::Max, atomicMax) -KMM_GPU_ATOMIC_REDUCTION_IMPL(unsigned long long int, Reduction::Max, atomicMax) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(__half, ReductionOp::Max, atomicMax) -//KMM_GPU_ATOMIC_REDUCTION_IMPL(__nv_bfloat16, ReductionOp::Max, atomicMax) - -template -struct IsGPUAtomicSupported: std::false_type {}; - -template -struct IsGPUAtomicSupported())>>: std::true_type {}; - -} // namespace kmm diff --git a/src/memops/gpu_reduction.cu b/src/memops/gpu_reduction.cu deleted file mode 100644 index 258f015b..00000000 --- a/src/memops/gpu_reduction.cu +++ /dev/null @@ -1,312 +0,0 @@ -#include - -#include "gpu_operators.cuh" - -#include "kmm/memops/gpu_fill.hpp" -#include "kmm/memops/gpu_reduction.hpp" -#include "kmm/utils/checked_math.hpp" -#include "kmm/utils/gpu_utils.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -static constexpr size_t total_block_size = 256; - -template -__global__ void reduction_kernel( - const T* src_buffer, - T* dst_buffer, - size_t num_outputs, - size_t num_inputs_per_output, - size_t input_stride, - size_t items_per_thread -) { - __shared__ T shared_results[total_block_size]; - - uint32_t thread_x = threadIdx.x; - uint32_t thread_y = threadIdx.y; - - uint64_t global_x = blockIdx.x * uint64_t(blockDim.x) + thread_x; - uint64_t global_y = blockIdx.y * uint64_t(blockDim.y) * items_per_thread + thread_y; - - ReductionOperator reduce; - T local_result = reduce.identity(); - - if (global_x < num_outputs && global_y < num_inputs_per_output) { - size_t x = global_x; - size_t max_y = min(global_y + items_per_thread * blockDim.y, num_inputs_per_output); - - for (size_t y = global_y; y < max_y; y += blockDim.y) { - T partial_result = src_buffer[y * input_stride + x]; - local_result = reduce(local_result, partial_result); - } - } - - if constexpr (UseSmem) { - shared_results[thread_y * blockDim.x + thread_x] = local_result; - - __syncthreads(); - - if (thread_y == 0) { - for (unsigned int y = 1; y < blockDim.y; y++) { - T partial_result = shared_results[y * blockDim.x + thread_x]; - local_result = reduce(local_result, partial_result); - } - } - } - - if (global_x < num_outputs && thread_y == 0) { - if constexpr (UseAtomics) { - GPUAtomic>::atomic_combine( - &dst_buffer[global_x], - local_result - ); - } else { - dst_buffer[global_x] = local_result; - } - } -} - -template -void execute_reduction_for_type_and_op( - g_stream_t stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - size_t num_outputs, - size_t num_partials_per_output, - size_t input_stride -) { - size_t block_size_x; - size_t block_size_y; - size_t items_per_thread; - - if (num_partials_per_output <= 8) { - block_size_x = total_block_size; - block_size_y = 1; - items_per_thread = num_partials_per_output; - } else if (num_outputs < 32) { - block_size_x = round_up_to_power_of_two(num_outputs); - block_size_y = total_block_size / block_size_x; - items_per_thread = 1; - } else { - block_size_x = 32; - block_size_y = total_block_size / block_size_x; - items_per_thread = 8; - } - - // Divide the total number of elements by the total number of threads - // We use max 512 blocks on the GPU as a rough heuristic here - size_t max_blocks_per_gpu = 512; - size_t max_grid_size_y = div_ceil(max_blocks_per_gpu, div_ceil(num_outputs, block_size_x)); - - // If we do not have atomics, we can only have 1 block in the y-direction - if (!IsGPUAtomicSupported>()) { - max_grid_size_y = 1; - } - - // The minimum items per thread is the number of partials divided by the maximum threads along Y - size_t min_items_per_thread = div_ceil(num_partials_per_output, max_grid_size_y * block_size_y); - - if (items_per_thread < min_items_per_thread) { - items_per_thread = min_items_per_thread; - } - - dim3 block_size = { - checked_cast(block_size_x), - checked_cast(block_size_y), - }; - - dim3 grid_size = { - checked_cast(div_ceil(num_outputs, block_size_x)), - checked_cast( - div_ceil(num_partials_per_output, block_size_y * items_per_thread) - ) - }; - - if constexpr (IsReductionSupported()) { - if (grid_size.y == 1 && block_size_y == 1) { - reduction_kernel<<>>( - reinterpret_cast(src_buffer), - reinterpret_cast(dst_buffer), - num_outputs, - num_partials_per_output, - input_stride, - items_per_thread - ); - - KMM_GPU_CHECK(gpu_get_last_error()); - return; - } - - if (grid_size.y == 1) { - reduction_kernel<<>>( - reinterpret_cast(src_buffer), - reinterpret_cast(dst_buffer), - num_outputs, - num_partials_per_output, - input_stride, - items_per_thread - ); - - KMM_GPU_CHECK(gpu_get_last_error()); - return; - } - - if constexpr (IsGPUAtomicSupported>()) { - T identity = ReductionOperator::identity(); - execute_gpu_fill_async(stream, dst_buffer, FillDef::with_value(identity, num_outputs)); - - reduction_kernel<<>>( - reinterpret_cast(src_buffer), - reinterpret_cast(dst_buffer), - num_outputs, - num_partials_per_output, - input_stride, - items_per_thread - ); - - KMM_GPU_CHECK(gpu_get_last_error()); - return; - } - } - - // silence unused warnings - (void)input_stride; - (void)stream; - (void)src_buffer; - (void)dst_buffer; - (void)block_size; - - throw std::runtime_error( - fmt::format("reduction {} for data type {} is not yet supported", Op, DataType::of()) - ); -} - -template -void execute_reduction_for_type( - g_stream_t stream, - Reduction operation, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - size_t num_outputs, - size_t num_partials_per_output, - size_t input_stride -) { -#define KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(O) \ - execute_reduction_for_type_and_op( \ - stream, \ - src_buffer, \ - dst_buffer, \ - num_outputs, \ - num_partials_per_output, \ - input_stride \ - ); - - switch (operation) { - case Reduction::Sum: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::Sum) - break; - case Reduction::Product: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::Product) - break; - case Reduction::Min: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::Min) - break; - case Reduction::Max: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::Max) - break; - case Reduction::BitAnd: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::BitAnd) - break; - case Reduction::BitOr: - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(Reduction::BitOr) - break; - default: - throw std::runtime_error( - fmt::format("reductions for operation {} are not yet supported", operation) - ); - } -} - -void execute_gpu_reduction_async( - g_stream_t stream, - g_device_ptr_t src_buffer, - g_device_ptr_t dst_buffer, - ReductionDef reduction -) { -#define KMM_CALL_REDUCTION_FOR_TYPE(T) \ - execute_reduction_for_type( \ - stream, \ - reduction.operation, \ - (g_device_ptr_t)((unsigned long long)(src_buffer) \ - + reduction.input_offset_elements * sizeof(T)), \ - (g_device_ptr_t)((unsigned long long)(dst_buffer) \ - + reduction.output_offset_elements * sizeof(T)), \ - reduction.num_outputs, \ - reduction.num_inputs_per_output, \ - reduction.input_stride_elements \ - ); - -#define KMM_CALL_REDUCTION_FOR_COMPLEX(T) \ - execute_reduction_for_type( \ - stream, \ - reduction.operation, \ - (g_device_ptr_t)((unsigned long long)(src_buffer) \ - + reduction.input_offset_elements * 2 * sizeof(T)), \ - (g_device_ptr_t)((unsigned long long)(dst_buffer) \ - + reduction.output_offset_elements * 2 * sizeof(T)), \ - 2 * reduction.num_outputs, \ - reduction.num_inputs_per_output, \ - 2 * reduction.input_stride_elements \ - ); - - switch (reduction.data_type.as_scalar()) { - case ScalarType::Int8: - KMM_CALL_REDUCTION_FOR_TYPE(int8_t) - return; - case ScalarType::Int16: - KMM_CALL_REDUCTION_FOR_TYPE(int16_t) - return; - case ScalarType::Int32: - KMM_CALL_REDUCTION_FOR_TYPE(int32_t) - return; - case ScalarType::Int64: - KMM_CALL_REDUCTION_FOR_TYPE(int64_t) - return; - case ScalarType::Uint8: - KMM_CALL_REDUCTION_FOR_TYPE(uint8_t) - return; - case ScalarType::Uint16: - KMM_CALL_REDUCTION_FOR_TYPE(uint16_t) - return; - case ScalarType::Uint32: - KMM_CALL_REDUCTION_FOR_TYPE(uint32_t) - return; - case ScalarType::Uint64: - KMM_CALL_REDUCTION_FOR_TYPE(uint64_t) - return; - case ScalarType::Float32: - KMM_CALL_REDUCTION_FOR_TYPE(float) - return; - case ScalarType::Float64: - KMM_CALL_REDUCTION_FOR_TYPE(double) - return; - case ScalarType::KeyAndInt64: - KMM_CALL_REDUCTION_FOR_TYPE(KeyValue) - return; - case ScalarType::KeyAndFloat64: - KMM_CALL_REDUCTION_FOR_TYPE(KeyValue) - return; - case ScalarType::Complex32: - KMM_CALL_REDUCTION_FOR_COMPLEX(float) - return; - case ScalarType::Complex64: - KMM_CALL_REDUCTION_FOR_COMPLEX(double) - return; - default: - throw std::runtime_error( - fmt::format("reductions on data type {} are not yet supported", reduction.data_type) - ); - } -} -} // namespace kmm \ No newline at end of file diff --git a/src/memops/host_copy.cpp b/src/memops/host_copy.cpp deleted file mode 100644 index f409c0a5..00000000 --- a/src/memops/host_copy.cpp +++ /dev/null @@ -1,60 +0,0 @@ -#include -#include - -#include "kmm/memops/host_copy.hpp" - -namespace kmm { - -template -inline void execute_copy_impl( - const void* src_buffer, - void* dst_buffer, - const CopyDef& copy_description, - I element_size -) { - for (size_t i2 = 0; i2 < copy_description.counts[2]; i2++) { - for (size_t i1 = 0; i1 < copy_description.counts[1]; i1++) { - for (size_t i0 = 0; i0 < copy_description.counts[0]; i0++) { - size_t src_offset = copy_description.src_offset - + (i0 * copy_description.src_strides[0]) - + (i1 * copy_description.src_strides[1]) - + (i2 * copy_description.src_strides[2]); - - size_t dst_offset = copy_description.dst_offset - + (i0 * copy_description.dst_strides[0]) - + (i1 * copy_description.dst_strides[1]) - + (i2 * copy_description.dst_strides[2]); - - ::memcpy( - static_cast(dst_buffer) + dst_offset, - static_cast(src_buffer) + src_offset, - element_size - ); - } - } - } -} - -template -bool is_aligned(const void* src_buffer, void* dst_buffer, const CopyDef& copy_description) { - bool result = reinterpret_cast(src_buffer) % Align == 0 - && reinterpret_cast(dst_buffer) % Align == 0 - && copy_description.src_offset % Align == 0 && copy_description.dst_offset % Align == 0 - && copy_description.element_size % Align == 0; - - for (size_t i = 0; i < CopyDef::MAX_DIMS; i++) { - if (copy_description.counts[i] > 1) { - result &= copy_description.src_strides[i] % Align == 0; - result &= copy_description.dst_strides[i] % Align == 0; - } - } - - return result; -} - -void execute_copy(const void* src_buffer, void* dst_buffer, CopyDef copy_def) { - copy_def.simplify(); - execute_copy_impl(src_buffer, dst_buffer, copy_def, copy_def.element_size); -} - -} // namespace kmm \ No newline at end of file diff --git a/src/memops/host_fill.cpp b/src/memops/host_fill.cpp deleted file mode 100644 index 520af0e0..00000000 --- a/src/memops/host_fill.cpp +++ /dev/null @@ -1,19 +0,0 @@ -#include - -#include "kmm/memops/host_fill.hpp" - -namespace kmm { - -void execute_fill(void* dst_buffer, const FillDef& fill) { - // TODO: optimize - size_t k = fill.fill_value.size(); - - for (size_t i = 0; i < fill.num_elements; i++) { - for (size_t j = 0; j < k; j++) { - static_cast(dst_buffer)[((fill.offset_elements + i) * k) + j] = - fill.fill_value[j]; - } - } -} - -} // namespace kmm diff --git a/src/memops/host_operators.hpp b/src/memops/host_operators.hpp deleted file mode 100644 index e5144331..00000000 --- a/src/memops/host_operators.hpp +++ /dev/null @@ -1,139 +0,0 @@ -#pragma once - -#include "kmm/core/reduction.hpp" -#include "kmm/utils/macros.hpp" - -namespace kmm { - -template -struct ReductionOperator; - -template -struct ReductionOperator< - T, - Reduction::Sum, - std::void_t() + std::declval())>> { - static KMM_HOST_DEVICE T identity() { - return T(0); - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return static_cast(a + b); - } -}; - -template -struct ReductionOperator< - T, - Reduction::Product, - std::void_t() * std::declval())>> { - static KMM_HOST_DEVICE T identity() { - return T(1); - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return static_cast(a * b); - } -}; - -template -struct ReductionOperator::is_specialized>> { - static constexpr T MAX_VALUE = std::numeric_limits::max(); - - static KMM_HOST_DEVICE T identity() { - return MAX_VALUE; - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return a < b ? a : b; - } -}; - -template -struct ReductionOperator::is_specialized>> { - static constexpr T MIN_VALUE = std::numeric_limits::lowest(); - - static KMM_HOST_DEVICE T identity() { - return MIN_VALUE; - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return b < a ? a : b; - } -}; - -template -struct ReductionOperator>> { - static KMM_HOST_DEVICE T identity() { - return T(0); - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return static_cast(a | b); - } -}; - -template -struct ReductionOperator>> { - static KMM_HOST_DEVICE T identity() { - // Note: we need the static cast here since decltype(~(short)0) == int - return static_cast(~T(0)); - } - - KMM_HOST_DEVICE T operator()(T a, T b) { - return static_cast(a & b); - } -}; - -template<> -struct ReductionOperator { - static KMM_HOST_DEVICE float identity() { - return INFINITY; - } - - KMM_HOST_DEVICE float operator()(float a, float b) { - return fminf(a, b); - } -}; - -template<> -struct ReductionOperator { - static KMM_HOST_DEVICE float identity() { - return -INFINITY; - } - - KMM_HOST_DEVICE float operator()(float a, float b) { - return fmaxf(a, b); - } -}; - -template<> -struct ReductionOperator { - static KMM_HOST_DEVICE double identity() { - return -double(INFINITY); - } - - KMM_HOST_DEVICE double operator()(double a, double b) { - return fmax(a, b); - } -}; - -template<> -struct ReductionOperator { - static KMM_HOST_DEVICE double identity() { - return double(INFINITY); - } - - KMM_HOST_DEVICE double operator()(double a, double b) { - return fmin(a, b); - } -}; - -template -struct IsReductionSupported: std::false_type {}; - -template -struct IsReductionSupported())>>: - std::true_type {}; - -} // namespace kmm \ No newline at end of file diff --git a/src/memops/host_reduction.cpp b/src/memops/host_reduction.cpp deleted file mode 100644 index 0ebd5a52..00000000 --- a/src/memops/host_reduction.cpp +++ /dev/null @@ -1,163 +0,0 @@ -#include "host_operators.hpp" - -#include "kmm/memops/host_reduction.hpp" - -namespace kmm { - -template -KMM_NOINLINE void execute_reduction_few_rows( - const T* __restrict__ src_buffer, - T* __restrict__ dst_buffer, - size_t num_columns, - size_t row_stride -) { - for (size_t j = 0; j < num_columns; j++) { - dst_buffer[j] = src_buffer[j]; - - for (size_t i = 1; i < NumRows; i++) { - dst_buffer[j] = - ReductionOperator()(dst_buffer[j], src_buffer[(i * row_stride) + j]); - } - } -} - -template -KMM_NOINLINE void execute_reduction_basic( - const T* __restrict__ src_buffer, - T* __restrict__ dst_buffer, - size_t num_columns, - size_t num_rows, - size_t row_stride -) { - for (size_t j = 0; j < num_columns; j++) { - dst_buffer[j] = src_buffer[j]; - } - - for (size_t i = 1; i < num_rows; i++) { - for (size_t j = 0; j < num_columns; j++) { - dst_buffer[j] = - ReductionOperator()(dst_buffer[j], src_buffer[(i * row_stride) + j]); - } - } -} - -template -KMM_NOINLINE void execute_reduction_impl( - const T* src_buffer, - T* dst_buffer, - size_t num_columns, - size_t num_rows, - size_t row_stride -) { - // For zero rows, we just fill with the identity value. - if (num_rows == 0) { - std::fill_n(dst_buffer, num_columns, ReductionOperator::identity()); - return; - } - -#define KMM_IMPL_REDUCTION_CASE(N) \ - if (num_rows == (N)) { \ - return execute_reduction_few_rows<(N), T, Op>( \ - src_buffer, \ - dst_buffer, \ - num_columns, \ - row_stride \ - ); \ - } - - // Specialize based on the number of rows - KMM_IMPL_REDUCTION_CASE(1) - KMM_IMPL_REDUCTION_CASE(2) - KMM_IMPL_REDUCTION_CASE(3) - KMM_IMPL_REDUCTION_CASE(4) - KMM_IMPL_REDUCTION_CASE(5) - KMM_IMPL_REDUCTION_CASE(6) - KMM_IMPL_REDUCTION_CASE(7) - KMM_IMPL_REDUCTION_CASE(8) - - return execute_reduction_basic( - src_buffer, - dst_buffer, - num_columns, - num_rows, - row_stride - ); -} - -void execute_reduction(const void* src_buffer, void* dst_buffer, ReductionDef reduction) { -#define KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, OP) \ - if constexpr (IsReductionSupported()) { \ - if (reduction.operation == Reduction::OP) { \ - execute_reduction_impl< \ - T, \ - Reduction::OP>(/* NOLINTNEXTLINE */ \ - static_cast(src_buffer) \ - + reduction.input_offset_elements, /* NOLINTNEXTLINE */ \ - static_cast(dst_buffer) + reduction.output_offset_elements, \ - reduction.num_outputs, \ - reduction.num_inputs_per_output, \ - reduction.input_stride_elements \ - ); \ - return; \ - } \ - } - -#define KMM_CALL_REDUCTION_FOR_TYPE(T) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, Sum) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, Product) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, Min) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, Max) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, BitAnd) \ - KMM_CALL_REDUCTION_FOR_TYPE_AND_OP(T, BitOr) - - switch (reduction.data_type.as_scalar()) { - case ScalarType::Int8: - KMM_CALL_REDUCTION_FOR_TYPE(int8_t) - break; - case ScalarType::Int16: - KMM_CALL_REDUCTION_FOR_TYPE(int16_t) - break; - case ScalarType::Int32: - KMM_CALL_REDUCTION_FOR_TYPE(int32_t) - break; - case ScalarType::Int64: - KMM_CALL_REDUCTION_FOR_TYPE(int64_t) - break; - case ScalarType::Uint8: - KMM_CALL_REDUCTION_FOR_TYPE(uint8_t) - break; - case ScalarType::Uint16: - KMM_CALL_REDUCTION_FOR_TYPE(uint16_t) - break; - case ScalarType::Uint32: - KMM_CALL_REDUCTION_FOR_TYPE(uint32_t) - break; - case ScalarType::Uint64: - KMM_CALL_REDUCTION_FOR_TYPE(uint64_t) - break; - case ScalarType::Float32: - KMM_CALL_REDUCTION_FOR_TYPE(float) - break; - case ScalarType::Float64: - KMM_CALL_REDUCTION_FOR_TYPE(double) - break; - case ScalarType::Complex32: - KMM_CALL_REDUCTION_FOR_TYPE(std::complex) - break; - case ScalarType::Complex64: - KMM_CALL_REDUCTION_FOR_TYPE(std::complex) - break; - case ScalarType::KeyAndInt64: - KMM_CALL_REDUCTION_FOR_TYPE(KeyValue) - break; - case ScalarType::KeyAndFloat64: - KMM_CALL_REDUCTION_FOR_TYPE(KeyValue) - break; - default: - break; - } - - throw std::runtime_error("unsupported reduction operation"); -} - -} // namespace kmm \ No newline at end of file diff --git a/src/memops/types.cpp b/src/memops/types.cpp deleted file mode 100644 index 3dd275ca..00000000 --- a/src/memops/types.cpp +++ /dev/null @@ -1,247 +0,0 @@ -#include -#include -#include - -#include "host_operators.hpp" - -#include "kmm/memops/types.hpp" -#include "kmm/utils/checked_math.hpp" - -namespace kmm { - -size_t CopyDef::minimum_source_bytes_needed() const { - size_t result = src_offset; - - for (size_t i = 0; i < MAX_DIMS; i++) { - if (counts[i] < 1) { - return 0; - } - - result += checked_mul(counts[i] - 1, src_strides[i]); - } - - return result + element_size; -} - -size_t CopyDef::minimum_destination_bytes_needed() const { - size_t result = dst_offset; - - for (size_t i = 0; i < MAX_DIMS; i++) { - if (counts[i] < 1) { - return 0; - } - - result += checked_mul(counts[i] - 1, dst_strides[i]); - } - - return result + element_size; -} - -void CopyDef::add_dimension(size_t count, size_t src_offset, size_t dst_offset) { - add_dimension(count, src_offset, dst_offset, 1, 1); -} - -void CopyDef::add_dimension( - size_t count, - size_t src_offset, - size_t dst_offset, - size_t src_stride, - size_t dst_stride -) { - this->src_offset += src_offset * src_stride; - this->dst_offset += dst_offset * dst_stride; - - if (src_stride == element_size && dst_stride == element_size) { - element_size *= count; - return; - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - if (src_stride == counts[i] * src_strides[i] && dst_stride == counts[i] * dst_strides[i]) { - counts[i] *= count; - return; - } - - if (counts[i] == 1) { - counts[i] = count; - src_strides[i] = src_stride; - dst_strides[i] = dst_stride; - return; - } - } - - throw std::length_error("the number of dimensions of a copy operation cannot exceed 3"); -} - -size_t CopyDef::effective_dimensionality() const { - for (size_t n = MAX_DIMS; n > 0; n--) { - if (counts[n - 1] != 1) { - return n; - } - } - - return 0; -} - -size_t CopyDef::number_of_bytes_copied() const { - return checked_mul(checked_product(counts, counts + MAX_DIMS), element_size); -} - -void CopyDef::simplify() { - if (number_of_bytes_copied() == 0) { - element_size = 0; - src_offset = 0; - dst_offset = 0; - - for (size_t i = 0; i < MAX_DIMS; i++) { - counts[i] = 0; - src_strides[i] = 1; - dst_strides[i] = 1; - } - - return; - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - for (size_t j = 0; j < MAX_DIMS; j++) { - if (src_strides[j] == element_size && dst_strides[j] == element_size) { - element_size *= counts[j]; - counts[j] = 1; - src_strides[j] = 1; - dst_strides[j] = 1; - } - } - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - for (size_t j = 0; j < MAX_DIMS; j++) { - if (i != j && src_strides[j] == counts[i] * src_strides[i] - && dst_strides[j] == counts[i] * dst_strides[i]) { - counts[i] *= counts[j]; - - counts[j] = 1; - src_strides[j] = 1; - dst_strides[j] = 1; - } - } - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - if (counts[i] == 1) { - src_strides[i] = 0; - dst_strides[i] = 0; - } - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - for (size_t j = i + 1; j < MAX_DIMS; j++) { - if ((counts[i] == 1 && counts[j] != 1) || dst_strides[i] > dst_strides[j] - || (dst_strides[i] == dst_strides[j] && src_strides[i] > src_strides[j])) { - std::swap(counts[i], counts[j]); - std::swap(src_strides[i], src_strides[j]); - std::swap(dst_strides[i], dst_strides[j]); - } - } - } - - for (size_t i = 0; i < MAX_DIMS; i++) { - if (counts[i] == 1) { - if (i == 0) { - src_strides[0] = element_size; - dst_strides[0] = element_size; - } else { - src_strides[i] = src_strides[i - 1]; - dst_strides[i] = dst_strides[i - 1]; - } - } - } -} - -size_t FillDef::minimum_destination_bytes_needed() const { - return checked_mul(checked_add(offset_elements, num_elements), fill_value.size()); -} - -size_t ReductionDef::minimum_destination_bytes_needed() const { - return checked_mul(data_type.size_in_bytes(), output_offset_elements + num_outputs); -} - -size_t ReductionDef::minimum_source_bytes_needed() const { - return checked_mul( - data_type.size_in_bytes(), - input_offset_elements + checked_mul(num_inputs_per_output, input_stride_elements) - ); -} - -[[noreturn]] void throw_invalid_reduction_exception(DataType dtype, Reduction op) { - throw std::runtime_error(fmt::format("invalid reduction {} for type {}", op, dtype)); -} - -template -std::vector identity_value_for_type_and_op() { - if constexpr (IsReductionSupported()) { - T value = ReductionOperator::identity(); - - uint8_t buffer[sizeof(T)]; - ::memcpy(buffer, &value, sizeof(T)); - return {buffer, buffer + sizeof(T)}; - } else { - throw_invalid_reduction_exception(DataType::of(), Op); - } -} - -template -std::vector identity_value_for_type(Reduction op) { - switch (op) { - case Reduction::Sum: - return identity_value_for_type_and_op(); - case Reduction::Product: - return identity_value_for_type_and_op(); - case Reduction::Min: - return identity_value_for_type_and_op(); - case Reduction::Max: - return identity_value_for_type_and_op(); - case Reduction::BitAnd: - return identity_value_for_type_and_op(); - case Reduction::BitOr: - return identity_value_for_type_and_op(); - default: - throw_invalid_reduction_exception(DataType::of(), op); - } -} - -std::vector reduction_identity_value(DataType dtype, Reduction op) { - switch (dtype.as_scalar()) { - case ScalarType::Int8: - return identity_value_for_type(op); - case ScalarType::Int16: - return identity_value_for_type(op); - case ScalarType::Int32: - return identity_value_for_type(op); - case ScalarType::Int64: - return identity_value_for_type(op); - case ScalarType::Uint8: - return identity_value_for_type(op); - case ScalarType::Uint16: - return identity_value_for_type(op); - case ScalarType::Uint32: - return identity_value_for_type(op); - case ScalarType::Uint64: - return identity_value_for_type(op); - case ScalarType::Float32: - return identity_value_for_type(op); - case ScalarType::Float64: - return identity_value_for_type(op); - case ScalarType::Complex32: - return identity_value_for_type>(op); - case ScalarType::Complex64: - return identity_value_for_type>(op); - case ScalarType::KeyAndInt64: - return identity_value_for_type>(op); - case ScalarType::KeyAndFloat64: - return identity_value_for_type>(op); - default: - throw_invalid_reduction_exception(dtype, op); - } -} - -} // namespace kmm \ No newline at end of file diff --git a/src/planner/array_descriptor.cpp b/src/planner/array_descriptor.cpp deleted file mode 100644 index f7ae1b8f..00000000 --- a/src/planner/array_descriptor.cpp +++ /dev/null @@ -1,226 +0,0 @@ -#include "kmm/memops/host_copy.hpp" -#include "kmm/planner/array_descriptor.hpp" -#include "kmm/planner/read_planner.hpp" -#include "kmm/planner/write_planner.hpp" -#include "kmm/runtime/runtime.hpp" - -namespace kmm { - -template -ArrayDescriptor::ArrayDescriptor( - TaskGraph& stage, - Distribution distribution, - DataType dtype -) : - m_distribution(std::move(distribution)), - m_dtype(dtype) { - size_t num_chunks = m_distribution.num_chunks(); - m_buffers.resize(num_chunks); - - for (size_t i = 0; i < num_chunks; i++) { - auto chunk = m_distribution.chunk(i); - auto num_elements = chunk.size.volume(); - auto layout = BufferLayout::for_type(dtype, num_elements); - auto buffer_id = stage.create_buffer(layout); - - m_buffers[i] = BufferDescriptor {.id = buffer_id, .layout = layout}; - } -} - -template -class CopyIntoTask: public ComputeTask { - public: - CopyIntoTask(void* dst_buffer, CopyDef copy_def) { - m_dst_buffer = dst_buffer; - m_copy = copy_def; - } - - void execute(Resource& resource, TaskContext context) { - KMM_ASSERT(context.accessors.size() == 1); - const void* src_buffer = context.accessors[0].address; - execute_copy(src_buffer, m_dst_buffer, m_copy); - } - - private: - void* m_dst_buffer; - CopyDef m_copy; -}; - -template -class CopyFromTask: public ComputeTask { - public: - CopyFromTask(const void* src_buffer, CopyDef copy_def) { - m_src_buffer = src_buffer; - m_copy = copy_def; - } - - void execute(Resource& resource, TaskContext context) { - KMM_ASSERT(context.accessors.size() == 1); - KMM_ASSERT(context.accessors[0].is_writable); - void* dst_buffer = context.accessors[0].address; - execute_copy(m_src_buffer, dst_buffer, m_copy); - } - - private: - const void* m_src_buffer; - CopyDef m_copy; -}; - -template -CopyDef build_copy_operation( - Point src_offset, - Dim src_dims, - Point dst_offset, - Dim dst_dims, - Dim counts, - size_t element_size -) { - auto copy_def = CopyDef(element_size); - - size_t src_stride = element_size; - size_t dst_stride = element_size; - - for (size_t i = 0; is_less(i, N); i++) { - copy_def.add_dimension( // - checked_cast(counts[i]), - checked_cast(src_offset[i]), - checked_cast(dst_offset[i]), - src_stride, - dst_stride - ); - - src_stride *= checked_cast(src_dims[i]); - dst_stride *= checked_cast(dst_dims[i]); - } - - copy_def.simplify(); - return copy_def; -} - -template -EventId ArrayDescriptor::copy_bytes_into_buffer(TaskGraph& stage, void* dst_data) { - std::unique_lock guard(m_mutex); - - auto& dist = this->distribution(); - auto data_type = this->data_type(); - auto num_chunks = dist.num_chunks(); - auto new_read_events = EventList {}; - - for (size_t i = 0; i < num_chunks; i++) { - auto& buffer = m_buffers[i]; - - auto buffer_req = BufferRequirement { - .buffer_id = buffer.id, - .memory_id = MemoryId::host(), - .access_mode = AccessMode::Read - }; - - auto chunk = dist.chunk(i); - auto chunk_region = Bounds::from_offset_size(chunk.offset, chunk.size); - - auto copy_def = build_copy_operation( // - Point::zero(), - chunk_region.size(), - chunk_region.begin(), - dist.array_size(), - chunk_region.size(), - data_type.size_in_bytes() - ); - - auto task = std::make_unique>(dst_data, copy_def); - - auto event_id = stage.insert_compute_task( - ResourceId::host(), - std::move(task), - {buffer_req}, - {buffer.last_write_event} - ); - - new_read_events.push_back(event_id); - } - - for (size_t i = 0; i < num_chunks; i++) { - m_buffers[i].last_access_events.push_back(new_read_events[i]); - } - - return stage.join_events(new_read_events); -} - -template -EventId ArrayDescriptor::copy_bytes_from_buffer(TaskGraph& stage, const void* src_data) { - std::shared_lock guard(m_mutex); - - auto& dist = this->distribution(); - auto data_type = this->data_type(); - auto num_chunks = dist.num_chunks(); - auto new_write_events = EventList {}; - - for (size_t i = 0; i < num_chunks; i++) { - auto& buffer = m_buffers[i]; - - auto buffer_req = BufferRequirement { - .buffer_id = buffer.id, // - .memory_id = MemoryId::host(), - .access_mode = AccessMode::ReadWrite - }; - - auto chunk = dist.chunk(i); - auto chunk_region = Bounds::from_offset_size(chunk.offset, chunk.size); - - auto copy_def = build_copy_operation( // - Point::zero(), - chunk_region.size(), - chunk_region.begin(), - dist.array_size(), - chunk_region.size(), - data_type.size_in_bytes() - ); - - auto task = std::make_unique>(src_data, copy_def); - - auto event_id = stage.insert_compute_task( - ResourceId::host(), - std::move(task), - {buffer_req}, - buffer.last_access_events - ); - - new_write_events.push_back(event_id); - } - - for (size_t i = 0; i < num_chunks; i++) { - auto& buffer = m_buffers[i]; - buffer.last_write_event = new_write_events[i]; - buffer.last_access_events = {new_write_events[i]}; - } - - return stage.join_events(new_write_events); -} - -template -EventId ArrayDescriptor::join_events(TaskGraph& stage) const { - std::unique_lock guard(m_mutex); - EventList deps; - - for (const auto& buffer : m_buffers) { - deps.insert_all(buffer.last_access_events); - } - - return stage.join_events(deps); -} - -template -void ArrayDescriptor::destroy(TaskGraph& stage) { - std::unique_lock guard(m_mutex); - - for (BufferDescriptor& meta : m_buffers) { - stage.delete_buffer(meta.id, std::move(meta.last_access_events)); - } - - m_distribution = Distribution(); - m_buffers.clear(); -} - -KMM_INSTANTIATE_ARRAY_IMPL(ArrayDescriptor) - -} // namespace kmm diff --git a/src/planner/read_planner.cpp b/src/planner/read_planner.cpp deleted file mode 100644 index 135b2ffe..00000000 --- a/src/planner/read_planner.cpp +++ /dev/null @@ -1,64 +0,0 @@ -#include "kmm/planner/read_planner.hpp" - -namespace kmm { - -template -ArrayReadPlanner::ArrayReadPlanner(std::shared_ptr> instance) : - m_lock(instance->m_mutex, std::try_to_lock), - m_instance(std::move(instance)) { - KMM_ASSERT(m_instance); - - if (!m_lock) { - throw std::runtime_error( - "array could not be locked for writing, which may happen if the " - "same array is also written to by the same kernel" - ); - } -} - -template -ArrayReadPlanner::~ArrayReadPlanner() {} - -template -BufferRequirement ArrayReadPlanner::prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - EventList& deps_out -) { - KMM_ASSERT(m_instance); - size_t chunk_index = m_instance->m_distribution.region_to_chunk_index(region); - auto chunk = m_instance->m_distribution.chunk(chunk_index); - const auto& buffer = m_instance->m_buffers[chunk_index]; - - region = Bounds::from_offset_size(chunk.offset, chunk.size); - deps_out.push_back(buffer.last_write_event); - m_read_events.push_back({chunk_index, EventId()}); - - return BufferRequirement { - .buffer_id = buffer.id, - .memory_id = memory_id, - .access_mode = AccessMode::Read - }; -} - -template -void ArrayReadPlanner::finalize_access(TaskGraph& stage, EventId event_id) { - KMM_ASSERT(!m_read_events.empty()); - m_read_events.back().second = event_id; -} - -template -void ArrayReadPlanner::commit(TaskGraph& stage) { - KMM_ASSERT(m_instance); - - for (const auto& [chunk_index, event_id] : m_read_events) { - m_instance->m_buffers[chunk_index].last_access_events.push_back(event_id); - } - - m_read_events.clear(); -} - -KMM_INSTANTIATE_ARRAY_IMPL(ArrayReadPlanner) - -} // namespace kmm \ No newline at end of file diff --git a/src/planner/reduction_planner.cpp b/src/planner/reduction_planner.cpp deleted file mode 100644 index fc4dbe85..00000000 --- a/src/planner/reduction_planner.cpp +++ /dev/null @@ -1,256 +0,0 @@ -#include "kmm/planner/reduction_planner.hpp" -#include "kmm/runtime/task_graph.hpp" - -namespace kmm { - -template -ArrayReductionPlanner::ArrayReductionPlanner( - std::shared_ptr> instance, - Reduction op -) : - m_lock(instance->m_mutex, std::try_to_lock), - m_instance(std::move(instance)), - m_reduction(op) { - KMM_ASSERT(m_instance); - - if (!m_lock) { - throw std::runtime_error( - "array could not be locked for reductions, which may happen if " - "the same array is provided multiple times as an argument to a kernel" - ); - } -} - -template -ArrayReductionPlanner::~ArrayReductionPlanner() {} - -template -BufferRequirement ArrayReductionPlanner::prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - size_t replication_factor, - EventList& deps_out -) { - size_t chunk_index = m_instance->m_distribution.region_to_chunk_index(region); - auto chunk = m_instance->m_distribution.chunk(chunk_index); - - region = Bounds::from_offset_size(chunk.offset, chunk.size); - - auto dtype = m_instance->data_type(); - auto num_elements = checked_mul(checked_cast(region.volume()), replication_factor); - BufferLayout layout = BufferLayout::for_type(dtype, num_elements); - - auto buffer_id = stage.create_buffer(layout); - - auto fill_event = stage.insert_node(CommandFill { - .dst_buffer = buffer_id, - .memory_id = memory_id, - .definition = FillDef( - dtype.size_in_bytes(), - num_elements, - reduction_identity_value(dtype, m_reduction).data() - ) - }); - - m_partial_buffers.push_back(PartialReductionBuffer { - .chunk_index = chunk_index, - .buffer_id = buffer_id, - .memory_id = memory_id, - .replication_factor = replication_factor, - .creation_event = fill_event, - .write_events = {} - }); - - deps_out.push_back(fill_event); - - return BufferRequirement { - .buffer_id = buffer_id, - .memory_id = memory_id, - .access_mode = AccessMode::Exclusive - }; -} - -template -void ArrayReductionPlanner::finalize_access(TaskGraph& stage, EventId event_id) { - KMM_ASSERT(!m_partial_buffers.empty()); - m_partial_buffers.back().write_events.push_back(event_id); -} - -template -std::pair ArrayReductionPlanner::reduce_per_chunk_and_memory( - TaskGraph& stage, - size_t chunk_index, - MemoryId memory_id, - PartialReductionBuffer** buffers, - size_t num_buffers -) { - auto dtype = m_instance->data_type(); - auto chunk = m_instance->distribution().chunk(chunk_index); - auto num_elements = checked_cast(chunk.size.volume()); - auto layout = BufferLayout::for_type(dtype, num_elements); - - auto scratch_buffer = stage.create_buffer(layout.repeat(num_buffers)); - auto scratch_writes = EventList {}; - - for (size_t i = 0; i < num_buffers; i++) { - EventId event_id = stage.insert_node( - CommandReduction { - buffers[i]->buffer_id, - scratch_buffer, - memory_id, - ReductionDef { - .operation = m_reduction, - .data_type = dtype, - .num_outputs = num_elements, - .num_inputs_per_output = buffers[i]->replication_factor, - .output_offset_elements = i * num_elements - }, - }, - std::move(buffers[i]->write_events) - ); - - stage.delete_buffer(buffers[i]->buffer_id, {event_id}); - scratch_writes.push_back(event_id); - } - - BufferId final_buffer = stage.create_buffer(layout); - EventId final_write = stage.insert_node( - CommandReduction { - scratch_buffer, - final_buffer, - memory_id, - ReductionDef { - .operation = m_reduction, - .data_type = dtype, - .num_outputs = num_elements, - .num_inputs_per_output = num_buffers, - }, - }, - std::move(scratch_writes) - ); - - stage.delete_buffer(scratch_buffer, {final_write}); - return {final_buffer, final_write}; -} - -template -EventId ArrayReductionPlanner::reduce_per_chunk( - TaskGraph& stage, - size_t chunk_index, - PartialReductionBuffer** buffers, - size_t num_buffers -) { - auto chunk = m_instance->distribution().chunk(chunk_index); - auto num_elements = checked_cast(chunk.size.volume()); - auto dtype = m_instance->data_type(); - auto layout = BufferLayout::for_type(dtype, num_elements); - - std::sort(buffers, buffers + num_buffers, [&](const auto* a, const auto* b) { - return a->memory_id < b->memory_id; - }); - - std::vector> intermediates; - - for (size_t begin = 0, end = begin; begin < num_buffers; begin = end) { - auto memory_id = buffers[begin]->memory_id; - - while (end < num_buffers && buffers[end]->memory_id == memory_id) { - end++; - } - - auto [buffer_id, event_id] = reduce_per_chunk_and_memory( - stage, - chunk_index, - memory_id, - &buffers[begin], - end - begin - ); - - intermediates.emplace_back(memory_id, buffer_id, event_id); - } - - auto collect_buffer = stage.create_buffer(layout.repeat(intermediates.size())); - auto collect_events = EventList {}; - - for (size_t i = 0; i < intermediates.size(); i++) { - auto [memory_id, buffer_id, event_id] = intermediates[i]; - - auto element_size = dtype.size_in_bytes(); - auto copy_definition = CopyDef(element_size); - copy_definition.add_dimension( - num_elements, // - 0, - i * num_elements, - element_size, - element_size - ); - - auto copy_event = stage.insert_node( - CommandCopy { - .src_buffer = buffer_id, // - .src_memory = memory_id, - .dst_buffer = collect_buffer, - .dst_memory = chunk.owner_id, - .definition = copy_definition - }, - {event_id} - ); - - collect_events.push_back(copy_event); - stage.delete_buffer(buffer_id, {copy_event}); - } - - auto& final_buffer = m_instance->m_buffers[chunk_index]; - collect_events.insert_all(final_buffer.last_access_events); - - auto final_event = stage.insert_node( - CommandReduction { - .src_buffer = collect_buffer, - .dst_buffer = final_buffer.id, - .memory_id = chunk.owner_id, - .definition = - ReductionDef { - .operation = m_reduction, - .data_type = m_instance->data_type(), - .num_outputs = num_elements, - .num_inputs_per_output = intermediates.size() - } - }, - std::move(collect_events) - ); - - stage.delete_buffer(collect_buffer, {final_event}); - - final_buffer.last_write_event = final_event; - final_buffer.last_access_events = {final_event}; - return final_event; -} - -template -void ArrayReductionPlanner::commit(TaskGraph& stage) { - auto buffers = std::vector(); - for (size_t i = 0; i < m_partial_buffers.size(); i++) { - buffers.push_back(&m_partial_buffers[i]); - } - - std::sort(buffers.begin(), buffers.end(), [&](const auto* a, const auto* b) { - return a->chunk_index < b->chunk_index; - }); - - for (size_t begin = 0, end = begin; begin < buffers.size(); begin = end) { - size_t chunk_index = buffers[begin]->chunk_index; - - while (end < buffers.size() && chunk_index == buffers[end]->chunk_index) { - end++; - } - - reduce_per_chunk(stage, chunk_index, &buffers[begin], end - begin); - } - - m_partial_buffers.clear(); -} - -KMM_INSTANTIATE_ARRAY_IMPL(ArrayReductionPlanner) - -} // namespace kmm \ No newline at end of file diff --git a/src/planner/write_planner.cpp b/src/planner/write_planner.cpp deleted file mode 100644 index 2bb44d23..00000000 --- a/src/planner/write_planner.cpp +++ /dev/null @@ -1,89 +0,0 @@ -#include - -#include "kmm/planner/write_planner.hpp" -#include "kmm/runtime/task_graph.hpp" - -namespace kmm { - -template -ArrayWritePlanner::ArrayWritePlanner(std::shared_ptr> instance) : - m_lock(instance->m_mutex, std::try_to_lock), - m_instance(std::move(instance)) { - KMM_ASSERT(m_instance); - - if (!m_lock) { - throw std::runtime_error( - "array could not be locked for writing, which may happen if the " - "same array is provided multiple times as an argument to a kernel" - ); - } -} - -template -ArrayWritePlanner::~ArrayWritePlanner() {} - -template -BufferRequirement ArrayWritePlanner::prepare_access( - TaskGraph& stage, - MemoryId memory_id, - Bounds& region, - EventList& deps_out -) { - size_t chunk_index = m_instance->m_distribution.region_to_chunk_index(region); - auto chunk = m_instance->m_distribution.chunk(chunk_index); - const auto& buffer = m_instance->m_buffers[chunk_index]; - - region = Bounds::from_offset_size(chunk.offset, chunk.size); - deps_out.insert_all(buffer.last_access_events); - m_write_events.push_back({chunk_index, EventId()}); - - return BufferRequirement { - .buffer_id = buffer.id, - .memory_id = memory_id, - .access_mode = AccessMode::ReadWrite - }; -} - -template -void ArrayWritePlanner::finalize_access(TaskGraph& stage, EventId event_id) { - KMM_ASSERT(!m_write_events.empty()); - m_write_events.back().second = event_id; -} - -template -void ArrayWritePlanner::commit(TaskGraph& stage) { - std::sort(m_write_events.begin(), m_write_events.end(), [&](const auto& a, const auto& b) { - return a.first < b.first; - }); - - for (size_t begin = 0, end = 0; begin < m_write_events.size(); begin = end) { - size_t chunk_index = m_write_events[begin].first; - auto& buffer = m_instance->m_buffers[chunk_index]; - - EventId write_event; - while (end < m_write_events.size() && chunk_index == m_write_events[end].first) { - end++; - } - - if (begin + 1 == end) { - write_event = m_write_events[begin].second; - } else { - EventList deps; - - for (size_t i = begin; i < end; i++) { - deps.push_back(m_write_events[i].second); - } - - write_event = stage.join_events(std::move(deps)); - } - - buffer.last_write_event = write_event; - buffer.last_access_events = {write_event}; - } - - m_write_events.clear(); -} - -KMM_INSTANTIATE_ARRAY_IMPL(ArrayWritePlanner) - -} // namespace kmm \ No newline at end of file diff --git a/src/runtime/allocators/arena.cpp b/src/runtime/allocators/arena.cpp new file mode 100644 index 00000000..1131363a --- /dev/null +++ b/src/runtime/allocators/arena.cpp @@ -0,0 +1,290 @@ +#include +#include + +#include "kmm/core/panic.hpp" +#include "kmm/runtime/allocators/arena.hpp" + +namespace kmm { + +static constexpr size_t ARENA_ALIGNMENT = 256; + +static size_t round_up_to_alignment(size_t n) { + return (n + ARENA_ALIGNMENT - 1) / ARENA_ALIGNMENT * ARENA_ALIGNMENT; +} + +ArenaAllocator::ArenaAllocator(std::unique_ptr inner, size_t block_size) : + m_inner(std::move(inner)), + m_block_size(block_size) {} + +ArenaAllocator::~ArenaAllocator() { + poll(); + + for (auto& block : m_blocks) { + m_inner->deallocate(block->base, BufferLayout {block->size}); + } +} + +void ArenaAllocator::insert_free(Block& block, size_t offset, size_t size, DeviceEventSet deps) { + block.by_size.insert({size, offset}); + block.by_offset.emplace(offset, Chunk {size, std::move(deps)}); +} + +ArenaAllocator::Chunk ArenaAllocator::take_free(Block& block, size_t offset) { + auto it = block.by_offset.find(offset); + KMM_ASSERT(it != block.by_offset.end()); + + Chunk chunk = std::move(it->second); + block.by_offset.erase(it); + block.by_size.erase({chunk.size, offset}); + + return chunk; +} + +bool ArenaAllocator::find_best_fit( + const DeviceStream* stream_opt, + size_t nbytes, + Block*& block_out, + size_t& offset_out +) const { + // Case 1: prefer a free chunk whose last use is already guaranteed to precede `stream`, + // so reusing it needs no extra wait. This is common when the same stream is reused. + size_t num_blocks = m_blocks.size(); + + for (size_t k = 0; k < num_blocks; k++) { + size_t block_index = (m_search_hint + k) % num_blocks; + const Block& block = *m_blocks[block_index]; + + for (auto it = block.by_size.lower_bound({nbytes, 0}); it != block.by_size.end(); ++it) { + const auto& deps = block.by_offset.at(it->second).deps; + bool is_ready = + stream_opt == nullptr ? m_events.is_ready(deps) : stream_opt->preceded_by(deps); + + if (is_ready) { + block_out = m_blocks[block_index].get(); + offset_out = it->second; + m_search_hint = block_index; + return true; + } + } + } + + // Case 2: no ready chunk exists anywhere, so fall back to the smallest fitting chunk. + for (size_t k = 0; k < num_blocks; k++) { + size_t block_index = (m_search_hint + k) % num_blocks; + auto it = m_blocks[block_index]->by_size.lower_bound({nbytes, 0}); + + if (it == m_blocks[block_index]->by_size.end()) { + continue; + } + + block_out = m_blocks[block_index].get(); + offset_out = it->second; + m_search_hint = block_index; + return true; + } + + return false; +} + +AllocResult ArenaAllocator::add_block(const DeviceStream* stream_opt, size_t min_size) { + min_size = std::max(min_size, size_t(1024)); + size_t size = std::max(m_block_size, min_size); + void* base; + + while (true) { + AllocResult result; + DeviceEvent event {}; + + if (stream_opt != nullptr) { + result = m_inner->allocate_async(*stream_opt, BufferLayout {size}, &base); + } else { + result = m_inner->allocate(BufferLayout {size}, &base); + } + + if (result == AllocResult::Success) { + if (stream_opt != nullptr) { + event = stream_opt->record_event(); + } + + auto block = std::make_unique(); + block->base = base; + block->size = size; + insert_free(*block, 0, size, event); + m_blocks.push_back(std::move(block)); + m_bytes_reserved += size; + return AllocResult::Success; + } + + // if we drop below the request size, we give up. + if (size <= min_size) { + return result; + } + + size = std::max(min_size, size / 2); + } +} + +AllocResult ArenaAllocator::allocate_generic( + const DeviceStream* stream_opt, + BufferLayout layout, + void** addr_out +) { + size_t nbytes = round_up_to_alignment(layout.size_in_bytes); + + Block* block; + size_t offset; + + while (!find_best_fit(stream_opt, nbytes, block, offset)) { + // did not find a block, try to add a block + AllocResult result = add_block(stream_opt, nbytes); + + // new block added, we have our best fit + if (result == AllocResult::Success) { + offset = 0; + block = m_blocks.back().get(); + break; + } + + // Could not find a block, try to deallocate an empty block. There might be + // enough memory available, just that the blocks are too fragmented. + if (trim_one(stream_opt)) { + continue; + } + + // could not find a block, could not allocate a new block, could not free an unused block. + // I am out of ideas. Just return that the allocation has failed. + return result; + } + + Chunk chunk = take_free(*block, offset); + + if (chunk.size > nbytes) { + insert_free(*block, offset + nbytes, chunk.size - nbytes, chunk.deps); + } + + if (stream_opt != nullptr) { + stream_opt->wait_on_event(chunk.deps); + } else { + m_events.synchronize(chunk.deps); + } + + *addr_out = static_cast(block->base) + offset; + m_allocations[*addr_out] = {block, offset, nbytes}; + return AllocResult::Success; +} + +void ArenaAllocator::deallocate_generic( + const DeviceStream* stream_opt, + void* addr, + BufferLayout layout +) { + auto it = m_allocations.find(addr); + KMM_ASSERT(it != m_allocations.end()); + + Allocation alloc = it->second; + m_allocations.erase(it); + + Block& block = *alloc.block; + size_t offset = alloc.offset; + size_t size = alloc.size; + DeviceEventSet deps; + + if (stream_opt != nullptr) { + deps.insert(stream_opt->record_event()); + } + + // Coalesce with the preceding free chunk + auto succ_it = block.by_offset.lower_bound(offset); + if (succ_it != block.by_offset.begin()) { + auto pred_it = std::prev(succ_it); + + if (pred_it->first + pred_it->second.size == offset) { + size_t pred_offset = pred_it->first; + Chunk pred_chunk = take_free(block, pred_offset); + + offset = pred_offset; + size += pred_chunk.size; + deps.insert(std::move(pred_chunk.deps)); + } + } + + // Coalesce with the next free chunk + auto next_it = block.by_offset.find(offset + size); + if (next_it != block.by_offset.end()) { + size_t next_offset = next_it->first; + Chunk next_chunk = take_free(block, next_offset); + + size += next_chunk.size; + deps.insert(std::move(next_chunk.deps)); + } + + insert_free(block, offset, size, std::move(deps)); +} + +AllocResult ArenaAllocator::allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out +) { + return allocate_generic(&stream, layout, addr_out); +} + +void ArenaAllocator::deallocate_async(const DeviceStream& stream, void* addr, BufferLayout layout) { + deallocate_generic(&stream, addr, layout); +} + +AllocResult ArenaAllocator::allocate(BufferLayout layout, void** addr_out) { + return allocate_generic(nullptr, layout, addr_out); +} + +void ArenaAllocator::deallocate(void* addr, BufferLayout layout) { + deallocate_generic(nullptr, addr, layout); +} + +void ArenaAllocator::poll() { + m_inner->poll(); +} + +void ArenaAllocator::trim(size_t nbytes_remaining) { + while (true) { + // we can quit, enough bytes are available + if (m_bytes_reserved <= nbytes_remaining) { + break; + } + + // try to release on unused block. + if (!trim_one(nullptr)) { + break; + } + } + + m_inner->trim(nbytes_remaining); +} + +bool ArenaAllocator::trim_one(const DeviceStream* stream_opt) { + for (size_t i = 0; i < m_blocks.size(); i++) { + Block& block = *m_blocks[i]; + const auto& [offset, chunk] = *block.by_offset.begin(); + bool fully_free = block.by_offset.size() == 1 && offset == 0 && chunk.size == block.size; + + if (fully_free) { + const auto& event = chunk.deps; + + if (stream_opt != nullptr) { + stream_opt->wait_on_event(event); + m_inner->deallocate_async(*stream_opt, block.base, BufferLayout {block.size}); + } else { + m_events.synchronize(event); + m_inner->deallocate(block.base, BufferLayout {block.size}); + } + + m_bytes_reserved -= block.size; + m_blocks.erase(m_blocks.begin() + i); + return true; + } + } + + return false; +} + +} // namespace kmm diff --git a/src/runtime/allocators/base.cpp b/src/runtime/allocators/base.cpp index 98527f7b..ad5cfbcc 100644 --- a/src/runtime/allocators/base.cpp +++ b/src/runtime/allocators/base.cpp @@ -1,84 +1,36 @@ +#include + #include "kmm/runtime/allocators/base.hpp" +#include "kmm/runtime/device_event.hpp" namespace kmm { -SyncAllocator::SyncAllocator(std::shared_ptr streams, size_t max_bytes) : - m_streams(streams), - m_bytes_limit(max_bytes), - m_bytes_in_use(0) {} - -SyncAllocator::~SyncAllocator() {} - -AllocationResult SyncAllocator::allocate_async( - size_t nbytes, - void** addr_out, - DeviceEventSet& deps_out -) { - KMM_ASSERT(nbytes > 0); - make_progress(); - - while (true) { - if (m_bytes_limit - m_bytes_in_use >= nbytes) { - auto result = this->allocate(nbytes, addr_out); - - if (result == AllocationResult::Success) { - m_bytes_in_use += nbytes; - return AllocationResult::Success; - } - } - - if (m_pending_deallocs.empty()) { - return AllocationResult::ErrorOutOfMemory; - } - - auto d = m_pending_deallocs.front(); - m_streams->wait_until_ready(d.dependencies); - m_pending_deallocs.pop_front(); - m_bytes_in_use -= d.nbytes; - - this->deallocate(d.addr, d.nbytes); - } -} - -void SyncAllocator::deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) { - make_progress(); +Allocator::Allocator() = default; - if (m_streams->is_ready(deps)) { - m_bytes_in_use -= nbytes; - this->deallocate(addr, nbytes); - } else { - m_pending_deallocs.push_back({addr, nbytes, std::move(deps)}); - } +Allocator::~Allocator() { + // Force through any deallocations still waiting on their dependencies, + // otherwise their memory would never actually be freed. + trim(0); } -void SyncAllocator::make_progress() { - while (!m_pending_deallocs.empty()) { - auto d = m_pending_deallocs.front(); - - if (!m_streams->is_ready(d.dependencies)) { - break; - } - - m_pending_deallocs.pop_front(); - - m_bytes_in_use -= d.nbytes; - this->deallocate(d.addr, d.nbytes); - } +AllocResult Allocator::allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out +) { + // we must wait for all events on the stream to finish + stream.synchronize(); + return allocate(layout, addr_out); } -void SyncAllocator::trim(size_t nbytes_remaining) { - while (m_bytes_in_use > nbytes_remaining) { - if (m_pending_deallocs.empty()) { - break; - } - - auto d = m_pending_deallocs.front(); - m_pending_deallocs.pop_front(); - - m_streams->wait_until_ready(d.dependencies); - m_bytes_in_use -= d.nbytes; - this->deallocate(d.addr, d.nbytes); - } +void Allocator::deallocate_async( // + const DeviceStream& stream, + void* addr, + BufferLayout layout +) { + // we must wait for all events on the stream to finish + stream.synchronize(); + deallocate(addr, layout); } -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/src/runtime/allocators/block.cpp b/src/runtime/allocators/block.cpp deleted file mode 100644 index b8865792..00000000 --- a/src/runtime/allocators/block.cpp +++ /dev/null @@ -1,284 +0,0 @@ -#include "kmm/runtime/allocators/block.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -static constexpr size_t MAX_ALIGNMENT = 256; - -struct BlockAllocator::Block { - std::set free_regions; - std::unique_ptr head = nullptr; - BlockRegion* tail = nullptr; - void* base_addr = nullptr; - size_t size = 0; - - Block(void* addr, size_t size, DeviceEventSet deps) { - this->head = std::make_unique(this, 0, size, (deps)); - this->tail = head.get(); - this->free_regions.insert(this->head.get()); - this->base_addr = addr; - this->size = size; - } -}; - -struct BlockAllocator::BlockRegion { - Block* parent = nullptr; - std::unique_ptr next = nullptr; - BlockRegion* prev = nullptr; - size_t offset_in_block = 0; - size_t size = 0; - bool is_free = false; - DeviceEventSet dependencies; - - BlockRegion(Block* parent, size_t offset_in_block, size_t size, DeviceEventSet deps = {}) { - this->parent = parent; - this->offset_in_block = offset_in_block; - this->size = size; - this->dependencies = (deps); - } -}; - -bool BlockAllocator::RegionSizeCompare::operator()(const BlockRegion* a, const BlockRegion* b) - const { - return a->size < b->size; -} - -bool BlockAllocator::RegionSizeCompare::operator()(const BlockRegion* a, BlockRegionSize b) const { - return a->size < b.value; -} - -BlockAllocator::BlockAllocator(std::unique_ptr allocator, size_t min_block_size) : - m_allocator(std::move(allocator)), - m_min_block_size(min_block_size) {} - -BlockAllocator::~BlockAllocator() { - for (size_t index = 0; index < m_blocks.size(); index++) { - auto& block = m_blocks[index]; - auto* region = m_blocks[index]->head.get(); - - if (!region->is_free || region->next != nullptr) { - // OH NO - continue; - } - - m_allocator->deallocate_async( // - block->base_addr, - block->size, - std::move(region->dependencies) - ); - } -} - -AllocationResult BlockAllocator::allocate_async( - size_t nbytes, - void** addr_out, - DeviceEventSet& deps_out -) { - size_t alignment = std::min(round_up_to_power_of_two(nbytes), MAX_ALIGNMENT); - nbytes = round_up_to_multiple(nbytes, alignment); - - auto* region = find_region(nbytes, alignment); - - if (region == nullptr) { - region = allocate_block(nbytes); - - if (region == nullptr) { - return AllocationResult::ErrorOutOfMemory; - } - } - - auto* block = region->parent; - block->free_regions.erase(region); - region->is_free = false; - - auto offset_in_region = offset_to_alignment(region, alignment); - - if (region->size > offset_in_region + nbytes) { - auto [left, right] = split_region(region, offset_in_region + nbytes); - block->free_regions.emplace(right); - region = left; - } - - *addr_out = static_cast(block->base_addr) + region->offset_in_block + offset_in_region; - deps_out.insert(region->dependencies); - - m_active_regions.emplace(addr_out, region); - return AllocationResult::Success; -} - -BlockAllocator::BlockRegion* BlockAllocator::allocate_block(size_t min_nbytes) { - DeviceEventSet deps; - void* base_addr; - size_t block_size = std::max(min_nbytes, m_min_block_size); - - while (true) { - auto result = m_allocator->allocate_async(block_size, &base_addr, deps); - - if (result == AllocationResult::Success) { - break; - } - - block_size /= 2; - - if (block_size < min_nbytes) { - return nullptr; - } - } - - auto new_block = std::make_unique(base_addr, block_size, std::move(deps)); - auto* region = new_block->head.get(); - - m_bytes_allocated += block_size; - m_blocks.insert( - m_blocks.begin() + static_cast(m_active_block), - std::move(new_block) - ); - - return region; -} - -BlockAllocator::BlockRegion* BlockAllocator::find_region(size_t nbytes, size_t alignment) { - for (size_t i = 0; i < m_blocks.size(); i++) { - auto* block = m_blocks[m_active_block].get(); - auto it = block->free_regions.lower_bound(BlockRegionSize {nbytes}); - - while (it != block->free_regions.end()) { - auto* region = &**it; - it++; - - if (fits_in_region(region, nbytes, alignment)) { - return region; - } - } - - m_active_block = (m_active_block + 1) % m_blocks.size(); - } - - return nullptr; -} - -void BlockAllocator::deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) { - auto it = m_active_regions.find(addr); - KMM_ASSERT(it != m_active_regions.end()); - - auto* region = it->second; - m_active_regions.erase(it); - - KMM_ASSERT(nbytes <= region->size); - KMM_ASSERT(region->is_free == false); - - region->is_free = true; - region->dependencies = std::move(deps); - - auto* block = region->parent; - auto* prev = region->prev; - auto* next = region->next.get(); - - if (prev != nullptr && prev->is_free) { - block->free_regions.erase(prev); - region = merge_regions(prev, region); - } - - if (next != nullptr && next->is_free) { - block->free_regions.erase(next); - region = merge_regions(region, next); - } - - block->free_regions.insert(region); -} - -size_t BlockAllocator::offset_to_alignment(const BlockRegion* region, size_t alignment) { - return round_up_to_multiple(region->offset_in_block, alignment) - region->offset_in_block; -} - -bool BlockAllocator::fits_in_region(const BlockRegion* region, size_t nbytes, size_t alignment) { - return region->size >= offset_to_alignment(region, alignment) + nbytes; -} - -auto BlockAllocator::split_region(BlockRegion* region, size_t left_size) - -> std::pair { - KMM_ASSERT(region->size > left_size); - - auto* parent = region->parent; - auto* left = region; - - size_t right_offset = left->offset_in_block + left_size; - size_t right_size = left->size - left_size; - left->size = left_size; - - auto right = - std::make_unique(parent, right_offset, right_size, left->dependencies); - auto* right_ptr = right.get(); - - if (left->next != nullptr) { - left->next->prev = right_ptr; - right->next = std::move(left->next); - } else { - parent->tail = right_ptr; - right->next = nullptr; - } - - right->prev = left; - left->next = std::move(right); - - return {left, right_ptr}; -} - -auto BlockAllocator::merge_regions(BlockRegion* left, BlockRegion* right) -> BlockRegion* { - auto* parent = left->parent; - - KMM_ASSERT(left->parent == parent && right->parent == parent); - KMM_ASSERT(left->is_free && right->is_free); - KMM_ASSERT(left->next.get() == right); - KMM_ASSERT(left == right->prev); - - if (right->next != nullptr) { - right->next->prev = left; - } else { - parent->tail = left; - } - - left->size += right->size; - left->dependencies.insert(right->dependencies); - left->next = std::move(right->next); // `right` is deleted here (since left.next == right) - return left; -} - -void BlockAllocator::make_progress() { - m_allocator->make_progress(); -} - -void BlockAllocator::trim(size_t nbytes_remaining) { - size_t index = 0; - - while (m_bytes_allocated >= nbytes_remaining) { - if (index >= m_blocks.size()) { - break; - } - - auto& block = m_blocks[index]; - auto* region = block->head.get(); - - if (!region->is_free || region->next != nullptr) { - index++; - continue; - } - - m_allocator->deallocate_async( // - block->base_addr, - block->size, - std::move(region->dependencies) - ); - - m_bytes_allocated -= block->size; - m_blocks.erase(m_blocks.begin() + static_cast(index)); - - if (m_active_block > index) { - m_active_block--; - } - } - - m_allocator->trim(nbytes_remaining); -} - -} // namespace kmm diff --git a/src/runtime/allocators/caching.cpp b/src/runtime/allocators/caching.cpp deleted file mode 100644 index f31fef9d..00000000 --- a/src/runtime/allocators/caching.cpp +++ /dev/null @@ -1,191 +0,0 @@ -#include "kmm/runtime/allocators/caching.hpp" -#include "kmm/utils/integer_fun.hpp" - -namespace kmm { - -CachingAllocator::CachingAllocator( - std::unique_ptr allocator, - double max_fragmentation, - size_t initial_watermark -) : - m_allocator(std::move(allocator)), - m_bytes_watermark(initial_watermark), - m_max_fragmentation(max_fragmentation){ - KMM_ASSERT(m_allocator != nullptr); - KMM_ASSERT(m_max_fragmentation >= 0.0 && m_max_fragmentation < 1.0); -} - -CachingAllocator::~CachingAllocator() { - while (free_some_memory() > 0) { - // - } -} - -struct CachingAllocator::AllocationSlot { - AllocationSlot(void* addr, size_t nbytes, DeviceEventSet dependencies) : - addr(addr), - nbytes(nbytes), - dependencies(std::move(dependencies)) {} - - void* addr = nullptr; - size_t nbytes = 0; - DeviceEventSet dependencies; - std::unique_ptr next = nullptr; - AllocationSlot* lru_older = nullptr; - AllocationSlot* lru_newer = nullptr; -}; - -size_t round_up_allocation_size(size_t nbytes) { - if (nbytes >= 1024) { - return round_up_to_multiple(nbytes, size_t(1024)); - } else { - return round_up_to_power_of_two(nbytes); - } -} - -bool CachingAllocator::can_allocate_bytes(size_t nbytes) const { - // If `m_bytes_allocated + nbytes <= m_bytes_watermark` then allocating nbytes will not - // raise the watermark, so it is allowed - if (nbytes <= m_bytes_watermark - m_bytes_allocated) { - return true; - } - - - // If `m_bytes_allocated - m_bytes_in_use == 0` then there is no overhead at all, so we - // must allow the allocation. - auto overhead = m_bytes_allocated - m_bytes_in_use; - if (overhead == 0) { - return true; - } - - // Otherwise, measure if the new overhead exceeds the maximum allowed fragmentation - auto new_watermark = m_bytes_allocated + nbytes; - return double(overhead) <= m_max_fragmentation * double(new_watermark); -} - -AllocationResult CachingAllocator::allocate_async( - size_t nbytes, - void** addr_out, - DeviceEventSet& deps_out -) { - nbytes = round_up_allocation_size(nbytes); - auto& bin = m_allocation_bins[nbytes]; - - if (bin.head == nullptr) { - while (true) { - if (can_allocate_bytes(nbytes)) { - auto result = m_allocator->allocate_async(nbytes, addr_out, deps_out); - - if (result == AllocationResult::Success) { - m_bytes_allocated += nbytes; - m_bytes_in_use += nbytes; - m_bytes_watermark = std::max(m_bytes_watermark, m_bytes_allocated); - return AllocationResult::Success; - } - } - - if (free_some_memory() == 0) { - return AllocationResult::ErrorOutOfMemory; - } - } - } - - auto slot = std::move(bin.head); - - if (slot->next != nullptr) { - bin.head = std::move(slot->next); - } else { - bin.tail = nullptr; - } - - if (slot->lru_older != nullptr) { - slot->lru_older->lru_newer = slot->lru_newer; - } else { - m_lru_oldest = slot->lru_newer; - } - - if (slot->lru_newer != nullptr) { - slot->lru_newer->lru_older = slot->lru_older; - } else { - m_lru_newest = slot->lru_older; - } - - m_bytes_in_use += nbytes; - *addr_out = slot->addr; - deps_out.insert(std::move(slot->dependencies)); - return AllocationResult::Success; -} - -void CachingAllocator::deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) { - nbytes = round_up_allocation_size(nbytes); - m_bytes_in_use -= nbytes; - - auto slot = std::make_unique(addr, nbytes, std::move(deps)); - auto* slot_addr = slot.get(); - - if (m_lru_newest != nullptr) { - m_lru_newest->lru_newer = slot_addr; - slot->lru_older = m_lru_newest; - m_lru_newest = slot_addr; - } else { - m_lru_newest = slot_addr; - m_lru_oldest = slot_addr; - } - - auto& bin = m_allocation_bins[nbytes]; - if (bin.head == nullptr) { - bin.head = std::move(slot); - bin.tail = slot_addr; - } else { - bin.tail->next = std::move(slot); - bin.tail = slot_addr; - } -} - -void CachingAllocator::make_progress() { - m_allocator->make_progress(); -} - -void CachingAllocator::trim(size_t nbytes_remaining) { - while (m_bytes_allocated > nbytes_remaining) { - if (free_some_memory() == 0) { - break; - } - } - - m_allocator->trim(nbytes_remaining); -} - -size_t CachingAllocator::free_some_memory() { - if (m_lru_oldest == nullptr) { - return 0; - } - - auto nbytes = m_lru_oldest->nbytes; - auto& bin = m_allocation_bins[nbytes]; - - KMM_ASSERT(bin.head.get() == m_lru_oldest); - auto slot = std::move(bin.head); - - if (slot->next != nullptr) { - bin.head = std::move(slot->next); - } else { - bin.tail = nullptr; - } - - KMM_ASSERT(slot->lru_older == nullptr); - - if (auto* newer = slot->lru_newer) { - newer->lru_older = nullptr; - m_lru_oldest = newer; - } else { - m_lru_oldest = nullptr; - m_lru_newest = nullptr; - } - - m_bytes_allocated -= slot->nbytes; - m_allocator->deallocate_async(slot->addr, slot->nbytes, std::move(slot->dependencies)); - return slot->nbytes; -} - -} // namespace kmm \ No newline at end of file diff --git a/src/runtime/allocators/device.cpp b/src/runtime/allocators/device.cpp index 6019bbca..e3cae5c6 100644 --- a/src/runtime/allocators/device.cpp +++ b/src/runtime/allocators/device.cpp @@ -1,196 +1,34 @@ -#include "kmm/runtime/allocators/device.hpp" - -namespace kmm { +#include -PinnedMemoryAllocator::PinnedMemoryAllocator( - GPUContextHandle context, - std::shared_ptr streams, - size_t max_bytes -) : - SyncAllocator(streams, max_bytes), - m_context(context) {} +#include "spdlog/spdlog.h" -AllocationResult PinnedMemoryAllocator::allocate(size_t nbytes, void** addr_out) { - GPUContextGuard guard {m_context}; - g_result_t result = - g_mem_host_alloc(addr_out, nbytes, G_MEMHOSTALLOC_PORTABLE | G_MEMHOSTALLOC_DEVICEMAP); - - if (result == G_SUCCESS) { - return AllocationResult::Success; - } else if (result == G_ERROR_OUT_OF_MEMORY) { - return AllocationResult::ErrorOutOfMemory; - } else { - throw GPUDriverException("error when calling `cuMemHostAlloc`", result); - } -} +#include "kmm/runtime/allocators/device.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" -void PinnedMemoryAllocator::deallocate(void* addr, size_t nbytes) { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(g_mem_free_host(addr)); -} +namespace kmm { -DeviceMemoryAllocator::DeviceMemoryAllocator( - GPUContextHandle context, - std::shared_ptr streams, - size_t max_bytes -) : - SyncAllocator(streams, max_bytes), - m_context(context) {} +DeviceMemoryAllocator::DeviceMemoryAllocator(g_context_t context) : m_context(context) {} -AllocationResult DeviceMemoryAllocator::allocate(size_t nbytes, void** addr_out) { +AllocResult DeviceMemoryAllocator::allocate(BufferLayout layout, void** addr_out) { GPUContextGuard guard {m_context}; g_device_ptr_t ptr; - g_result_t result = g_mem_alloc(&ptr, nbytes); + g_result_t result = g_mem_alloc(&ptr, layout.size_in_bytes); - if (result == G_SUCCESS) { - *addr_out = (void*)ptr; - return AllocationResult::Success; - } else if (result == G_ERROR_OUT_OF_MEMORY) { - return AllocationResult::ErrorOutOfMemory; - } else { - throw GPUDriverException("error when calling `g_mem_alloc`", result); + if (result == G_ERROR_OUT_OF_MEMORY) { + return AllocResult::ErrorOutOfMemory; } -} -void DeviceMemoryAllocator::deallocate(void* addr, size_t nbytes) { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(g_mem_free((g_device_ptr_t)addr)); + KMM_GPU_CHECK(result); + *addr_out = (void*)ptr; + spdlog::trace("allocate {} bytes of device memory (addr: {})", layout.size_in_bytes, *addr_out); + return AllocResult::Success; } -DevicePoolAllocator::DevicePoolAllocator( - GPUContextHandle context, - std::shared_ptr streams, - DevicePoolKind kind, - size_t max_bytes -) : - m_context(context), - m_streams(streams), - m_alloc_stream(streams->create_stream(context)), - m_dealloc_stream(streams->create_stream(context)), - m_kind(kind), - m_bytes_limit(max_bytes) { +void DeviceMemoryAllocator::deallocate(void* addr, BufferLayout layout) { + spdlog::trace("deallocate {} bytes of device memory (addr: {})", layout.size_in_bytes, addr); GPUContextGuard guard {m_context}; - - g_device_t device; - KMM_GPU_CHECK(g_ctx_get_device(&device)); - - switch (m_kind) { - case DevicePoolKind::Default: - KMM_GPU_CHECK(g_device_get_default_mem_pool(&m_pool, device)); - break; - - case DevicePoolKind::Create: - g_mem_pool_props_t props; - ::memset(&props, 0, sizeof(g_mem_pool_props_t)); - - props.allocType = G_MEM_ALLOCATION_TYPE_PINNED; - props.handleTypes = G_MEM_HANDLE_TYPE_NONE; - props.location.type = G_MEM_LOCATION_TYPE_DEVICE; - props.location.id = device; - - KMM_GPU_CHECK(g_mem_pool_create(&m_pool, &props)); - break; - } + KMM_GPU_CHECK(g_mem_free(g_device_ptr_t(addr))); } -DevicePoolAllocator::~DevicePoolAllocator() { - for (auto d : m_pending_deallocs) { - m_bytes_in_use -= d.nbytes; - m_streams->wait_until_ready(d.event); - } - - KMM_ASSERT(m_bytes_in_use == 0); - - GPUContextGuard guard {m_context}; - - switch (m_kind) { - case DevicePoolKind::Default: - // No need to destroy the default pool - break; - case DevicePoolKind::Create: - KMM_GPU_CHECK(g_mem_pool_destroy(m_pool)); - break; - } -} - -AllocationResult DevicePoolAllocator::allocate_async( - size_t nbytes, - void** addr_out, - DeviceEventSet& deps_out -) { - make_progress(); - - while (m_bytes_limit - m_bytes_in_use < nbytes) { - if (m_pending_deallocs.empty()) { - return AllocationResult::ErrorOutOfMemory; - } - - auto& d = m_pending_deallocs.front(); - m_streams->wait_for_event(m_alloc_stream, d.event); - m_bytes_in_use -= d.nbytes; - m_pending_deallocs.pop_front(); - } - - g_device_ptr_t device_ptr; - g_result_t result = g_result_t(G_ERROR_UNKNOWN); - - auto event = m_streams->with_stream(m_alloc_stream, [&](auto stream) { - GPUContextGuard guard {m_context}; - result = g_mem_alloc_from_pool_async(&device_ptr, nbytes, m_pool, stream); - }); - - if (result == G_SUCCESS) { - m_bytes_in_use += nbytes; - deps_out.insert(event); - *addr_out = (void*)device_ptr; - return AllocationResult::Success; - } else if (result == G_ERROR_OUT_OF_MEMORY) { - return AllocationResult::ErrorOutOfMemory; - } else { - throw GPUDriverException("error while calling `g_mem_alloc_from_pool_async`", result); - } -} - -void DevicePoolAllocator::deallocate_async(void* addr, size_t nbytes, DeviceEventSet deps) { - g_device_ptr_t device_ptr = (g_device_ptr_t)addr; - - auto event = m_streams->with_stream(m_dealloc_stream, deps, [&](auto stream) { - KMM_GPU_CHECK(g_mem_free_async(device_ptr, stream)); - }); - - m_pending_deallocs.push_back({.addr = addr, .nbytes = nbytes, .event = event}); -} - -void DevicePoolAllocator::make_progress() { - while (true) { - if (m_pending_deallocs.empty()) { - break; - } - - auto& d = m_pending_deallocs.front(); - - if (!m_streams->is_ready(d.event)) { - break; - } - - m_bytes_in_use -= d.nbytes; - m_pending_deallocs.pop_front(); - } -} - -void DevicePoolAllocator::trim(size_t nbytes_remaining) { - while (m_bytes_in_use > nbytes_remaining) { - if (m_pending_deallocs.empty()) { - break; - } - - auto& d = m_pending_deallocs.front(); - m_streams->wait_until_ready(d.event); - - m_bytes_in_use -= d.nbytes; - m_pending_deallocs.pop_front(); - } - - KMM_GPU_CHECK(g_mem_pool_trim_to(m_pool, nbytes_remaining)); -} -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/src/runtime/allocators/device_pool.cpp b/src/runtime/allocators/device_pool.cpp new file mode 100644 index 00000000..a3641d80 --- /dev/null +++ b/src/runtime/allocators/device_pool.cpp @@ -0,0 +1,154 @@ +#include +#include + +#include "spdlog/spdlog.h" + +#include "kmm/runtime/allocators/device_pool.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +DevicePoolAllocator::DevicePoolAllocator( + g_context_t context, + DevicePoolKind kind, + size_t max_size +) : + m_context(context), + m_pool(nullptr), + m_kind(kind) { + GPUContextGuard guard {m_context}; + + g_device_t device; + KMM_GPU_CHECK(g_ctx_get_device(&device)); + + // CUDA assumes maxSize is ignored if its zero, while this constructor uses max_size==MAX + if (max_size == std::numeric_limits::max()) { + max_size = 0; + } + + switch (m_kind) { + case DevicePoolKind::Default: + KMM_GPU_CHECK(g_device_get_default_mem_pool(&m_pool, device)); + break; + + case DevicePoolKind::Create: +#if defined(KMM_USE_CUDA) + CUmemPoolProps props; + ::memset(&props, 0, sizeof(CUmemPoolProps)); + + props.allocType = CUmemAllocationType::CU_MEM_ALLOCATION_TYPE_PINNED; + props.handleTypes = CUmemAllocationHandleType::CU_MEM_HANDLE_TYPE_NONE; + props.location.type = CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE; + props.maxSize = max_size; + props.location.id = device; + + KMM_GPU_CHECK(g_mem_pool_create(&m_pool, &props)); +#else + throw std::runtime_error("memory pool is only supported with CUDA backend"); +#endif + break; + } +} + +DevicePoolAllocator::~DevicePoolAllocator() { + GPUContextGuard guard {m_context}; + + switch (m_kind) { + case DevicePoolKind::Default: + // No need to destroy the default pool + break; + case DevicePoolKind::Create: + KMM_GPU_CHECK(g_mem_pool_destroy(m_pool)); + break; + } +} + +AllocResult DevicePoolAllocator::allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out +) { + g_device_ptr_t device_ptr; + g_result_t result; + + GPUContextGuard guard {m_context}; + result = g_mem_alloc_from_pool_async(&device_ptr, layout.size_in_bytes, m_pool, stream); + + if (result == G_ERROR_OUT_OF_MEMORY) { + return AllocResult::ErrorOutOfMemory; + } + + KMM_GPU_CHECK(result); + *addr_out = (void*)device_ptr; + spdlog::trace( + "allocate {} bytes of device memory on stream {} (addr: {})", + layout.size_in_bytes, + stream.id(), + *addr_out + ); + return AllocResult::Success; +} + +void DevicePoolAllocator::deallocate_async( + const DeviceStream& stream, + void* addr, + BufferLayout layout +) { + g_device_ptr_t device_ptr = (g_device_ptr_t)addr; + spdlog::trace( + "deallocate {} bytes of device memory on stream {} (addr: {})", + layout.size_in_bytes, + stream.id(), + addr + ); + + GPUContextGuard guard {m_context}; + KMM_GPU_CHECK(g_mem_free_async(device_ptr, stream)); +} + +AllocResult DevicePoolAllocator::allocate(BufferLayout layout, void** addr_out) { + g_device_ptr_t device_ptr; + g_result_t result; + + GPUContextGuard guard {m_context}; + + // Route through the pool (via the legacy default stream) rather than `gpuMemAlloc`, so + // this allocation is still subject to the pool's `maxSize` and gets reclaimed by `trim`. + result = g_mem_alloc_from_pool_async(&device_ptr, layout.size_in_bytes, m_pool, nullptr); + + if (result == G_ERROR_OUT_OF_MEMORY) { + return AllocResult::ErrorOutOfMemory; + } + + KMM_GPU_CHECK(result); + KMM_GPU_CHECK(g_stream_synchronize(nullptr)); + + spdlog::trace( + "allocate {} bytes of device memory on stream NULL (addr: {})", + layout.size_in_bytes, + *addr_out + ); + *addr_out = (void*)device_ptr; + return AllocResult::Success; +} + +void DevicePoolAllocator::deallocate(void* addr, BufferLayout layout) { + g_device_ptr_t device_ptr = (g_device_ptr_t)addr; + spdlog::trace( + "deallocate {} bytes of device memory on stream NULL (addr: {})", + layout.size_in_bytes, + addr + ); + + GPUContextGuard guard {m_context}; + KMM_GPU_CHECK(g_mem_free_async(device_ptr, nullptr)); + KMM_GPU_CHECK(g_stream_synchronize(nullptr)); +} + +void DevicePoolAllocator::trim(size_t nbytes_remaining) { + GPUContextGuard guard {m_context}; + KMM_GPU_CHECK(g_mem_pool_trim_to(m_pool, nbytes_remaining)); +} + +} // namespace kmm diff --git a/src/runtime/allocators/limit.cpp b/src/runtime/allocators/limit.cpp new file mode 100644 index 00000000..1fad4a58 --- /dev/null +++ b/src/runtime/allocators/limit.cpp @@ -0,0 +1,168 @@ +#include +#include + +#include "kmm/runtime/allocators/limit.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +LimitAllocator::LimitAllocator( + std::unique_ptr inner, + DeviceEventRegistry events, + size_t max_size +) : + m_inner(std::move(inner)), + m_events(std::move(events)), + m_bytes_limit(max_size), + m_bytes_active(0), + m_bytes_pending(0) {} + +LimitAllocator::~LimitAllocator() { + // wait until all evenst are done + while (!m_pending_deallocs.empty()) { + auto front = m_pending_deallocs.front(); + m_events.synchronize(front.event); + m_bytes_pending -= front.nbytes; + m_pending_deallocs.pop_front(); + } + + KMM_ASSERT(m_bytes_active == 0 && m_bytes_pending == 0); +} + +AllocResult LimitAllocator::allocate_async( + const DeviceStream& stream, + BufferLayout layout, + void** addr_out +) { + size_t nbytes = layout.size_in_bytes; + poll(); + + auto deps = DeviceEventSet {}; + + if (!ensure_enough_space(&stream, nbytes)) { + return AllocResult::ErrorOutOfMemory; + } + + // the stream must wait until reaching the barrier + stream.wait_on_event(m_limit_barrier); + + auto result = m_inner->allocate_async(stream, layout, addr_out); + + if (result != AllocResult::Success) { + return result; + } + + m_bytes_active += nbytes; + return AllocResult::Success; +} + +void LimitAllocator::deallocate_async(const DeviceStream& stream, void* addr, BufferLayout layout) { + m_inner->deallocate_async(stream, addr, layout); + auto event = stream.record_event(); + m_bytes_active -= layout.size_in_bytes; + m_bytes_pending += layout.size_in_bytes; + m_pending_deallocs.push_back({addr, layout.size_in_bytes, event}); +} + +AllocResult LimitAllocator::allocate(BufferLayout layout, void** addr_out) { + size_t nbytes = layout.size_in_bytes; + poll(); + + auto deps = DeviceEventSet {}; + + if (!ensure_enough_space(nullptr, nbytes)) { + return AllocResult::ErrorOutOfMemory; + } + + // the stream must wait until reaching the barrier + for (const auto& dep : m_limit_barrier) { + m_events.synchronize(dep); + } + + auto result = m_inner->allocate(layout, addr_out); + + if (result != AllocResult::Success) { + return result; + } + + m_bytes_active += nbytes; + return AllocResult::Success; +} + +void LimitAllocator::deallocate(void* addr, BufferLayout layout) { + m_inner->deallocate(addr, layout); + m_bytes_active -= layout.size_in_bytes; +} + +void LimitAllocator::poll() { + while (!m_pending_deallocs.empty()) { + auto front = m_pending_deallocs.front(); + + if (!m_events.is_ready(front.event)) { + break; + } + + m_bytes_pending -= front.nbytes; + m_pending_deallocs.pop_front(); + } + + m_limit_barrier.prune(m_events); + m_inner->poll(); +} + +void LimitAllocator::trim(size_t nbytes_remaining) { + while (m_bytes_active + m_bytes_pending > nbytes_remaining) { + if (m_pending_deallocs.empty()) { + break; + } + + auto& d = m_pending_deallocs.front(); + m_events.synchronize(d.event); + + m_bytes_pending -= d.nbytes; + m_pending_deallocs.pop_front(); + } + + m_inner->trim(nbytes_remaining); +} + +std::optional LimitAllocator::bytes_reserved() const { + // Prefer the inner allocator's real reservation; otherwise fall back to what this layer has + // handed out (allocated plus not-yet-reclaimed), which is what it caps. + if (auto inner = m_inner->bytes_reserved()) { + return inner; + } + + return m_bytes_active + m_bytes_pending; +} + +bool LimitAllocator::ensure_enough_space(const DeviceStream* stream, size_t nbytes) { + // we can allocate now, we are done + size_t remaining_bytes = m_bytes_limit - m_bytes_active - m_bytes_pending; + if (remaining_bytes >= nbytes) { + return true; + } + + // even freeing everything pending would not be enough, no point waiting on any of it + if (m_bytes_limit - m_bytes_active < nbytes) { + return false; + } + + auto it = m_pending_deallocs.begin(); + + while (it != m_pending_deallocs.end()) { + m_limit_barrier.insert(it->event); + m_bytes_pending -= it->nbytes; + remaining_bytes += it->nbytes; + it = m_pending_deallocs.erase(it); + + if (remaining_bytes >= nbytes) { + return true; + } + } + + return false; +} + +} // namespace kmm diff --git a/src/runtime/allocators/managed.cpp b/src/runtime/allocators/managed.cpp new file mode 100644 index 00000000..a35ea802 --- /dev/null +++ b/src/runtime/allocators/managed.cpp @@ -0,0 +1,39 @@ +#include + +#include "spdlog/spdlog.h" + +#include "kmm/runtime/allocators/managed.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +ManagedMemoryAllocator::ManagedMemoryAllocator(g_context_t context) : m_context(context) {} + +AllocResult ManagedMemoryAllocator::allocate(BufferLayout layout, void** addr_out) { + GPUContextGuard guard {m_context}; + g_device_ptr_t ptr; + g_result_t result = g_mem_alloc_managed(&ptr, layout.size_in_bytes, G_MEM_ATTACH_GLOBAL); + + if (result == G_ERROR_OUT_OF_MEMORY) { + return AllocResult::ErrorOutOfMemory; + } + + KMM_GPU_CHECK(result); + *addr_out = (void*)ptr; + spdlog::trace( + "allocate {} bytes of managed memory (addr: {})", + layout.size_in_bytes, + *addr_out + ); + return AllocResult::Success; +} + +void ManagedMemoryAllocator::deallocate(void* addr, BufferLayout layout) { + spdlog::trace("deallocate {} bytes of managed memory (addr: {})", layout.size_in_bytes, addr); + + GPUContextGuard guard {m_context}; + KMM_GPU_CHECK(g_mem_free(g_device_ptr_t(addr))); +} + +} // namespace kmm diff --git a/src/runtime/allocators/pinned.cpp b/src/runtime/allocators/pinned.cpp new file mode 100644 index 00000000..2409d9de --- /dev/null +++ b/src/runtime/allocators/pinned.cpp @@ -0,0 +1,37 @@ +#include + +#include "spdlog/spdlog.h" + +#include "kmm/runtime/allocators/pinned.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +PinnedMemoryAllocator::PinnedMemoryAllocator(g_context_t context) : m_context(context) {} + +AllocResult PinnedMemoryAllocator::allocate(BufferLayout layout, void** addr_out) { + GPUContextGuard guard {m_context}; + g_result_t result = g_mem_host_alloc( + addr_out, + layout.size_in_bytes, + G_MEMHOSTALLOC_PORTABLE | G_MEMHOSTALLOC_DEVICEMAP + ); + + if (result == G_ERROR_OUT_OF_MEMORY) { + return AllocResult::ErrorOutOfMemory; + } + + KMM_GPU_CHECK(result); + spdlog::trace("allocate {} bytes of pinned memory (addr: {})", layout.size_in_bytes, *addr_out); + return AllocResult::Success; +} + +void PinnedMemoryAllocator::deallocate(void* addr, BufferLayout layout) { + spdlog::trace("deallocate {} bytes of pinned memory (addr: {})", layout.size_in_bytes, addr); + + GPUContextGuard guard {m_context}; + KMM_GPU_CHECK(g_mem_free_host(addr)); +} + +} // namespace kmm diff --git a/src/runtime/allocators/system.cpp b/src/runtime/allocators/system.cpp index 0da90d8d..c1056596 100644 --- a/src/runtime/allocators/system.cpp +++ b/src/runtime/allocators/system.cpp @@ -1,13 +1,23 @@ +#include "spdlog/spdlog.h" + #include "kmm/runtime/allocators/system.hpp" namespace kmm { -AllocationResult SystemAllocator::allocate(size_t nbytes, void** addr_out) { - *addr_out = malloc(nbytes); - return *addr_out != nullptr ? AllocationResult::Success : AllocationResult::ErrorOutOfMemory; +AllocResult SystemAllocator::allocate(BufferLayout layout, void** addr_out) { + *addr_out = malloc(layout.size_in_bytes); + + if (*addr_out == nullptr) { + return AllocResult::ErrorOutOfMemory; + } + + spdlog::trace("allocate {} bytes of system memory (addr: {})", layout.size_in_bytes, *addr_out); + return AllocResult::Success; } -void SystemAllocator::deallocate(void* addr, size_t nbytes) { +void SystemAllocator::deallocate(void* addr, BufferLayout layout) { + spdlog::trace("deallocate {} bytes of system memory (addr: {})", layout.size_in_bytes, addr); free(addr); } + } // namespace kmm \ No newline at end of file diff --git a/src/runtime/buffer_registry.cpp b/src/runtime/buffer_registry.cpp deleted file mode 100644 index 18e65f91..00000000 --- a/src/runtime/buffer_registry.cpp +++ /dev/null @@ -1,151 +0,0 @@ -#include "fmt/format.h" -#include "spdlog/spdlog.h" - -#include "kmm/runtime/buffer_registry.hpp" - -namespace kmm { - -BufferRegistry::BufferRegistry(std::shared_ptr memory_manager) : - m_memory_manager(memory_manager) { - KMM_ASSERT(m_memory_manager); -} - -BufferId BufferRegistry::add(BufferId buffer_id, BufferLayout layout) { - auto [it, success] = m_buffers.emplace(buffer_id, BufferMeta {}); - - if (!success) { - throw std::runtime_error( - fmt::format("could not add buffer {}: buffer already exists", buffer_id) - ); - } - - auto buffer = m_memory_manager->create_buffer(layout, std::to_string(buffer_id)); - it->second.buffer = buffer; - - return buffer_id; -} - -void BufferRegistry::remove(BufferId buffer_id) { - auto it = m_buffers.find(buffer_id); - - if (it == m_buffers.end()) { - throw std::runtime_error( - fmt::format("could not remove buffer {}: buffer not found", buffer_id) - ); - } - - m_memory_manager->delete_buffer(it->second.buffer); - m_buffers.erase(it); -} - -std::shared_ptr BufferRegistry::get(BufferId id) { - auto it = m_buffers.find(id); - - // Buffer not found, ignore - if (it == m_buffers.end()) { - throw std::runtime_error(fmt::format("could not retrieve buffer {}: buffer not found", id)); - } - - auto& meta = it->second; - - // If poisoned, throw exception - if (meta.poison_reason_opt != nullptr) { - throw PoisonException(*meta.poison_reason_opt); - } - - return meta.buffer; -} - -void BufferRegistry::poison(BufferId id, PoisonException reason) { - auto it = m_buffers.find(id); - - // Buffer not found, ignore - if (it == m_buffers.end()) { - return; - } - - auto& meta = it->second; - - // Buffer already poisoned, ignore - if (meta.poison_reason_opt != nullptr) { - return; - } - - spdlog::warn("buffer {} was poisoned: {}", id, reason.what()); - meta.poison_reason_opt = std::make_unique(std::move(reason)); -} - -BufferRequestList BufferRegistry::create_requests(const std::vector& buffers) { - auto parent = m_memory_manager->create_transaction(); - auto requests = BufferRequestList {}; - - try { - for (const auto& r : buffers) { - auto buffer = this->get(r.buffer_id); - auto req = m_memory_manager->create_request(buffer, r.memory_id, r.access_mode, parent); - requests.push_back(BufferRequest {req}); - } - - return requests; - } catch (...) { - // Release the requests that have been created so far. - for (const auto& r : requests) { - m_memory_manager->release_request(r); - } - - throw; - } -} - -Poll BufferRegistry::poll_requests( - const BufferRequestList& requests, - DeviceEventSet& dependencies_out -) { - Poll result = Poll::Ready; - - for (const auto& req : requests) { - if (m_memory_manager->poll_request(*req, dependencies_out) != Poll::Ready) { - result = Poll::Pending; - } - } - - return result; -} - -std::vector BufferRegistry::access_requests(const BufferRequestList& requests) { - auto accessors = std::vector {}; - - for (const auto& req : requests) { - accessors.push_back(m_memory_manager->get_accessor(*req)); - } - - return accessors; -} - -void BufferRegistry::release_requests(BufferRequestList& requests, DeviceEvent event) { - for (auto& req : requests) { - m_memory_manager->release_request(req, event); - } - - requests.clear(); -} - -void BufferRegistry::poison_all( - const std::vector& buffers, - PoisonException reason -) { - for (const auto& r : buffers) { - if (r.access_mode != AccessMode::Read) { - this->poison(r.buffer_id, reason); - } - } -} - -PoisonException::PoisonException(const std::string& error) { - m_message = error; -} - -const char* PoisonException::what() const noexcept { - return m_message.c_str(); -} -} // namespace kmm diff --git a/src/runtime/data_interfaces/external.cpp b/src/runtime/data_interfaces/external.cpp new file mode 100644 index 00000000..12820730 --- /dev/null +++ b/src/runtime/data_interfaces/external.cpp @@ -0,0 +1,77 @@ +#include "fmt/format.h" + +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/external.hpp" + +namespace kmm { + +ExternalDataInterface::ExternalDataInterface(void* ptr, size_t size_in_bytes, MemoryId memory_id) : + m_ptr(ptr), + m_size_in_bytes(size_in_bytes), + m_memory_id(memory_id) { + KMM_ASSERT(ptr != nullptr); +} + +void ExternalDataInterface::check_memory_id(MemoryId memory_id) const { + if (memory_id != m_memory_id) { + throw std::runtime_error( + fmt::format( + "cannot access or copy buffer on memory {}: buffer is pinned to memory {}", + memory_id, + m_memory_id + ) + ); + } +} + +size_t ExternalDataInterface::size_in_bytes() const noexcept { + return m_size_in_bytes; +} + +AllocResult ExternalDataInterface::allocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out +) { + check_memory_id(memory_id); + return AllocResult::Success; +} + +void ExternalDataInterface::deallocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + check_memory_id(memory_id); +} + +void* ExternalDataInterface::address(MemoryId memory_id) const noexcept { + // `allocate` already rejected any other memory, so this is just a sanity check. + KMM_ASSERT(memory_id == m_memory_id); + return m_ptr; +} + +void ExternalDataInterface::copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps, + DeviceEventSet& deps_out +) { + check_memory_id(src); + check_memory_id(dst); + deps_out.insert(deps); +} + +bool ExternalDataInterface::is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst +) const noexcept { + return src == m_memory_id && dst == m_memory_id; +} + +} // namespace kmm diff --git a/src/runtime/data_interfaces/flat.cpp b/src/runtime/data_interfaces/flat.cpp new file mode 100644 index 00000000..9e57349d --- /dev/null +++ b/src/runtime/data_interfaces/flat.cpp @@ -0,0 +1,243 @@ +#include +#include + +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/flat.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memory_system.hpp" + +namespace kmm { + +static BufferLayout normalize_buffer_layout(BufferLayout layout) { + static constexpr size_t max_align = 128; + size_t k = std::min(std::max(layout.size_in_bytes, layout.alignment), max_align); + size_t align = round_up_to_power_of_two(k); + return {round_up_to_multiple(layout.size_in_bytes, align), align}; +} + +FlatDataInterface::FlatDataInterface(BufferLayout layout, FillValue fill_value) : + m_layout(normalize_buffer_layout(layout)), + m_fill_value(std::move(fill_value)) {} + +size_t FlatDataInterface::size_in_bytes() const noexcept { + return m_layout.size_in_bytes; +} + +AllocResult FlatDataInterface::allocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out +) { + if (memory_id.is_host()) { + KMM_ASSERT(m_host_ptr == nullptr); + return system.allocate_host(m_layout, &m_host_ptr, stream_hint, deps_out); + } else { + auto id = memory_id.as_device(); + auto& ptr = m_device_ptrs[id.get()]; + KMM_ASSERT(ptr == 0); + return system.allocate_device(id, m_layout, &ptr, stream_hint, deps_out); + } +} + +void FlatDataInterface::deallocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + if (memory_id.is_host()) { + KMM_ASSERT(m_host_ptr != nullptr); + system.deallocate_host(m_host_ptr, m_layout, stream_hint, deps); + m_host_ptr = nullptr; + } else { + auto id = memory_id.as_device(); + auto& ptr = m_device_ptrs[id.get()]; + KMM_ASSERT(ptr != 0); + + system.deallocate_device(id, ptr, m_layout, stream_hint, deps); + ptr = 0; + } +} + +void* FlatDataInterface::address(MemoryId memory_id) const noexcept { + if (memory_id.is_host()) { + return m_host_ptr; + } else { + return (void*)m_device_ptrs[memory_id.as_device().get()]; + } +} + +void FlatDataInterface::copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps, + DeviceEventSet& deps_out +) { + size_t nbytes = m_layout.size_in_bytes; + + if (src.is_host() && dst.is_device()) { + auto id = dst.as_device(); + auto event = system.copy_host_to_device( // + id, + m_host_ptr, + m_device_ptrs[id.get()], + nbytes, + stream_hint, + deps + ); + + deps_out.insert(event); + return; + } + + if (src.is_device() && dst.is_host()) { + auto id = src.as_device(); + + auto event = system.copy_device_to_host( // + id, + m_device_ptrs[id.get()], + m_host_ptr, + nbytes, + stream_hint, + deps + ); + + deps_out.insert(event); + return; + } + + if (src.is_device() && dst.is_device()) { + auto src_id = src.as_device(); + auto dst_id = dst.as_device(); + + auto event = system.copy_device_to_device( + src_id, + dst_id, + m_device_ptrs[src_id.get()], + m_device_ptrs[dst_id.get()], + nbytes, + stream_hint, + deps + ); + + deps_out.insert(event); + return; + } + + KMM_PANIC("cannot copy from host memory to host memory"); +} + +bool FlatDataInterface::is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst +) const noexcept { + return system.is_copy_supported(src, dst); +} + +std::future FlatDataInterface::initialize_host( + MemorySystem& system, + const DeviceEventSet& deps +) { + if (m_fill_value.length == 0) { + return {}; + } + + size_t element_size = m_fill_value.length; + size_t count = m_layout.size_in_bytes / element_size; + + FillDescription description(m_fill_value); + description.add_dimension( + static_cast(count), + static_cast(element_size) + ); + + return system.fill_host(m_host_ptr, description, deps); +} + +DeviceEvent FlatDataInterface::initialize_device( + MemorySystem& system, + DeviceId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + if (m_fill_value.length == 0) { + return DeviceEvent::null(); + } + + KMM_ASSERT(m_layout.size_in_bytes % m_fill_value.length); + size_t element_size = m_fill_value.length; + size_t count = m_layout.size_in_bytes / m_fill_value.length; + + FillDescription description(m_fill_value); + description.add_dimension( + static_cast(count), + static_cast(element_size) + ); + + return system.fill_device( + memory_id, + m_device_ptrs[memory_id.get()], + description, + stream_hint, + deps + ); +} + +AllocResult FlatDataInterface::allocate_and_copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEventSet& deps_out +) { + if (src.is_host() && dst.is_device()) { + auto id = dst.as_device(); + auto& ptr = m_device_ptrs[id.get()]; + KMM_ASSERT(ptr == 0); + + auto dep_out = DeviceEvent {}; + auto result = system.allocate_device_and_copy_from_host( // + id, + m_layout, + &m_device_ptrs[id.get()], + m_host_ptr, + stream_hint, + deps_in, + dep_out + ); + + deps_out.insert(dep_out); + return result; + } + + if (src.is_device() && dst.is_host()) { + DeviceEventSet deps; + auto id = src.as_device(); + KMM_ASSERT(m_host_ptr == nullptr); + + auto dep_out = DeviceEvent {}; + auto result = system.allocate_host_and_copy_from_device( + m_layout, + &m_host_ptr, + id, + m_device_ptrs[id.get()], + stream_hint, + deps_in, + dep_out + ); + + deps_out.insert(dep_out); + return result; + } + + // just forward to the default impl. + return DataInterface::allocate_and_copy(system, src, dst, stream_hint, deps_in, deps_out); +} + +} // namespace kmm diff --git a/src/runtime/data_interfaces/managed.cpp b/src/runtime/data_interfaces/managed.cpp new file mode 100644 index 00000000..36c06dee --- /dev/null +++ b/src/runtime/data_interfaces/managed.cpp @@ -0,0 +1,146 @@ +#include +#include + +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/managed.hpp" +#include "kmm/runtime/memory_system.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +static BufferLayout normalize_buffer_layout(BufferLayout layout) { + static constexpr size_t max_align = 128; + size_t k = std::min(std::max(layout.size_in_bytes, layout.alignment), max_align); + size_t align = round_up_to_power_of_two(k); + return {round_up_to_multiple(layout.size_in_bytes, align), align}; +} + +ManagedDataInterface::ManagedDataInterface(BufferLayout layout, FillValue fill_value) : + m_layout(normalize_buffer_layout(layout)), + m_fill_value(std::move(fill_value)) {} + +size_t ManagedDataInterface::size_in_bytes() const noexcept { + return m_layout.size_in_bytes; +} + +AllocResult ManagedDataInterface::allocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out +) { + if (m_refcount == 0) { + DeviceEventSet deps; + auto result = system.allocate_managed(m_layout, &m_ptr, stream_hint, deps); + + if (result != AllocResult::Success) { + return result; + } + + m_alloc_deps = std::move(deps); + } + + m_refcount++; + deps_out.insert(m_alloc_deps); + return AllocResult::Success; +} + +void ManagedDataInterface::deallocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + KMM_ASSERT(m_refcount > 0); + m_dealloc_deps.insert(deps); + + if (--m_refcount == 0) { + system.deallocate_managed(m_ptr, m_layout, stream_hint, m_dealloc_deps); + m_ptr = nullptr; + m_alloc_deps.clear(); + m_dealloc_deps.clear(); + } +} + +void* ManagedDataInterface::address(MemoryId memory_id) const noexcept { + return m_ptr; +} + +bool ManagedDataInterface::is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst +) const noexcept { + return true; +} + +void ManagedDataInterface::hint_access( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + system.prefetch_managed(memory_id, m_ptr, m_layout, stream_hint, deps); +} + +void ManagedDataInterface::copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps, + DeviceEventSet& deps_out +) { + deps_out.insert(deps); +} + +std::future ManagedDataInterface::initialize_host( + MemorySystem& system, + const DeviceEventSet& deps +) { + if (m_fill_value.length == 0) { + return {}; + } + + size_t element_size = m_fill_value.length; + size_t count = m_layout.size_in_bytes / element_size; + + FillDescription description(m_fill_value); + description.add_dimension( + static_cast(count), + static_cast(element_size) + ); + + return system.fill_host(m_ptr, description, deps); +} + +DeviceEvent ManagedDataInterface::initialize_device( + MemorySystem& system, + DeviceId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + if (m_fill_value.length == 0) { + return DeviceEvent::null(); + } + + size_t element_size = m_fill_value.length; + size_t count = m_layout.size_in_bytes / element_size; + + FillDescription description(m_fill_value); + description.add_dimension( + static_cast(count), + static_cast(element_size) + ); + + return system.fill_device( // + memory_id, + reinterpret_cast(m_ptr), + description, + stream_hint, + deps + ); +} + +} // namespace kmm diff --git a/src/runtime/data_interfaces/pinned.cpp b/src/runtime/data_interfaces/pinned.cpp new file mode 100644 index 00000000..f1c2c5c5 --- /dev/null +++ b/src/runtime/data_interfaces/pinned.cpp @@ -0,0 +1,102 @@ +#include +#include + +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/pinned.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memory_system.hpp" + +namespace kmm { + +static BufferLayout normalize_buffer_layout(BufferLayout layout) { + static constexpr size_t max_align = 128; + size_t k = std::min(std::max(layout.size_in_bytes, layout.alignment), max_align); + size_t align = round_up_to_power_of_two(k); + return {round_up_to_multiple(layout.size_in_bytes, align), align}; +} + +PinnedDataInterface::PinnedDataInterface(BufferLayout layout) : + m_layout(normalize_buffer_layout(layout)) {} + +size_t PinnedDataInterface::size_in_bytes() const noexcept { + return m_layout.size_in_bytes; +} + +AllocResult PinnedDataInterface::allocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out +) { + if (m_refcount == 0) { + DeviceEventSet deps; + auto result = system.allocate_host(m_layout, &m_host_ptr, stream_hint, deps); + + if (result != AllocResult::Success) { + return result; + } + + m_alloc_deps = std::move(deps); + } + + // Devices reach the single host allocation through a mapped pointer. Resolve here now so + // `allocate` throws the exception and address remains exception-free. + if (!memory_id.is_host()) { + auto device_id = memory_id.as_device(); + m_device_ptrs[device_id.get()] = system.translate_host_pointer(device_id, m_host_ptr); + } + + m_refcount++; + deps_out.insert(m_alloc_deps); + return AllocResult::Success; +} + +void PinnedDataInterface::deallocate( + MemorySystem& system, + MemoryId memory_id, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + KMM_ASSERT(m_refcount > 0); + m_dealloc_deps.insert(deps); + + if (--m_refcount == 0) { + system.deallocate_host(m_host_ptr, m_layout, stream_hint, m_dealloc_deps); + m_host_ptr = nullptr; + for (auto& ptr : m_device_ptrs) { + ptr = nullptr; + } + m_alloc_deps.clear(); + m_dealloc_deps.clear(); + } +} + +void* PinnedDataInterface::address(MemoryId memory_id) const noexcept { + if (memory_id.is_host()) { + return m_host_ptr; + } + + return m_device_ptrs[memory_id.as_device().get()]; +} + +bool PinnedDataInterface::is_copy_supported( + MemorySystem& system, + MemoryId src, + MemoryId dst +) const noexcept { + return system.is_copy_supported(src, dst); +} + +void PinnedDataInterface::copy( + MemorySystem& system, + MemoryId src, + MemoryId dst, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps, + DeviceEventSet& deps_out +) { + deps_out.insert(deps); +} + +} // namespace kmm diff --git a/src/runtime/device_data_streams.cpp b/src/runtime/device_data_streams.cpp new file mode 100644 index 00000000..6390ff88 --- /dev/null +++ b/src/runtime/device_data_streams.cpp @@ -0,0 +1,300 @@ +#include +#include +#include +#include +#include +#include + +#include "ankerl/unordered_dense.h" +#include "spdlog/spdlog.h" + +#include "kmm/core/panic.hpp" +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +static constexpr size_t MAX_KINDS = 3; + +struct StreamSlot { + StreamKind kind; + DeviceStreamId id; + GPUStreamOwner native_stream; + uint64_t estimated_finish_time = 0; + DeviceEventSet last_preds {}; + DeviceEvent last_event {}; + size_t active_users = 0; + uint64_t last_acquired = 0; + + StreamSlot(StreamKind kind, DeviceStreamId id, GPUStreamOwner native_stream) : + kind(kind), + id(id), + native_stream(std::move(native_stream)) {} +}; + +struct DeviceDataStreams::Impl { + std::array, MAX_KINDS>, MAX_DEVICES> streams_per_device; + ankerl::unordered_dense::map id_to_slot; + std::map estimated_finish_time; + DeviceEventRegistry events; + uint64_t acquire_counter = 0; + + Impl(DeviceEventRegistry events) : events(std::move(events)) {} +}; + +DeviceDataStreams::DeviceDataStreams( + const SystemInfo& info, + DeviceEventRegistry events, + size_t num_d2d_streams, + size_t num_h2d_streams, + size_t num_d2h_streams +) : + m_impl(std::make_unique(events)) { + for (size_t i = 0; i < info.num_devices(); i++) { + auto* context = info.device(DeviceId(i)).context(); + std::array options = { + std::make_tuple(StreamKind::DeviceToDevice, num_d2d_streams, "d2d"), + std::make_tuple(StreamKind::HostToDevice, num_h2d_streams, "h2d"), + std::make_tuple(StreamKind::DeviceToHost, num_d2h_streams, "d2h") + }; + + for (auto [kind, n, name] : options) { + auto& slots = m_impl->streams_per_device[i][size_t(kind)]; + + if (n <= 0) { + throw std::runtime_error("number of streams must be positive"); + } + + for (size_t j = 0; j < n; j++) { + auto label = fmt::format("gpu{}-{}-{}", i, name, j); + + auto stream = GPUStreamOwner {context}; + auto id = events.register_stream(stream, label); + slots.emplace_back(kind, id, std::move(stream)); + } + + for (auto& slot : slots) { + m_impl->id_to_slot[slot.id] = &slot; + } + } + } +} + +DeviceDataStreams::DeviceDataStreams(DeviceDataStreams&&) noexcept = default; + +DeviceDataStreams::~DeviceDataStreams() { + for (auto& [id, slot] : m_impl->id_to_slot) { + m_impl->events.unregister_stream(id); + } +} + +void DeviceDataStreams::make_progress() { + auto it = m_impl->estimated_finish_time.begin(); + + while (it != m_impl->estimated_finish_time.end()) { + // not ready, exit now + if (!m_impl->events.is_ready(it->first)) { + break; + } + + // erase the event + it = m_impl->estimated_finish_time.erase(it); + } + + for (auto& [id, slot] : m_impl->id_to_slot) { + if (m_impl->events.is_ready(slot->last_event) && slot->active_users == 0) { + slot->estimated_finish_time = 0; + slot->last_acquired = 0; + } + } +} + +static bool deps_equal_ignoring_stream( + const DeviceEventRegistry& events, + const DeviceEventSet& a, + const DeviceEventSet& b, + DeviceStreamId ignore_stream +) { + for (const auto& event : a) { + if (event.stream() == ignore_stream || events.is_ready(event)) { + continue; + } + + bool result = false; + + for (const auto& e : b) { + if (event.precedes(e)) { + result = true; + } + } + + if (!result) { + return false; + } + } + + for (const auto& event : b) { + if (event.stream() == ignore_stream || events.is_ready(event)) { + continue; + } + + bool result = false; + + for (const auto& e : a) { + if (event.precedes(e)) { + result = true; + } + } + + if (!result) { + return false; + } + } + + return true; +} + +DeviceStreamId DeviceDataStreams::acquire_stream( + DeviceId device_id, + StreamKind kind, + const DeviceEventSet& deps +) { + KMM_ASSERT(device_id.get() < MAX_DEVICES && size_t(kind) < MAX_KINDS); + auto& slots = m_impl->streams_per_device[device_id.get()][size_t(kind)]; + + uint64_t expected_start_time = 0; + + for (const auto& dep : deps) { + // `lower_bound` lands exactly on `dep` if it has a recorded prediction; otherwise it + // gives a starting point to scan forward for the next event on the same stream, whose + // finish-time estimate is a valid (if possibly loose) upper bound on `dep`'s. + auto it = m_impl->estimated_finish_time.lower_bound(dep); + + while (it != m_impl->estimated_finish_time.end() && it->first.stream() != dep.stream()) { + ++it; + } + + if (it == m_impl->estimated_finish_time.end() || it->first.stream() != dep.stream()) { + continue; + } + + // This event has already finished, we can remove it and ignore it. + if (m_impl->events.is_ready(it->first)) { + m_impl->estimated_finish_time.erase(it); + continue; + } + + expected_start_time = std::max(expected_start_time, it->second); + } + + StreamSlot* best_slot = nullptr; + std::tuple best_key; + + spdlog::trace( + "selecting stream for {} (kind={}, deps={}, expected_start={})", + device_id, + kind, + deps, + expected_start_time + ); + + for (auto& slot : slots) { + bool slot_has_affinity = deps.contains(slot.last_event); + uint64_t slot_start_time = std::max(slot.estimated_finish_time, expected_start_time); + uint64_t slot_slack = slot_start_time - slot.estimated_finish_time; + + // For H2D/D2H, the hardware has only one copy engine per direction, so separate + // streams don't get parallelism. If this slot saw the exact same dependencies + // last time, there's nothing to lose by reusing it + if (kind != StreamKind::DeviceToDevice && !slot_has_affinity) { + slot_has_affinity |= + deps_equal_ignoring_stream(m_impl->events, slot.last_preds, deps, slot.id); + } + + // Priority, in order: + // - fewest active users, + // - reuse the stream that produced the dependency, + // - earliest start time, + // - lowest idle time (previous finish time - next start time), + // - longest since last selected. + auto slot_key = std::make_tuple( + slot.active_users, + !slot_has_affinity, + slot_start_time, + slot_slack, + slot.last_acquired + ); + + spdlog::trace( + " - slot {}: last_event={}, users={}, affinity={}, start={}, slack={}", + slot.id, + slot.last_event, + slot.active_users, + slot_has_affinity, + slot_start_time, + slot_slack + ); + + if (best_slot == nullptr || slot_key < best_key) { + best_slot = &slot; + best_key = slot_key; + } + } + + KMM_ASSERT(best_slot); + best_slot->active_users++; + best_slot->last_preds = deps; + best_slot->estimated_finish_time = + std::max(best_slot->estimated_finish_time, expected_start_time); + best_slot->last_acquired = ++m_impl->acquire_counter; + + // let the stream wait on the dependencies. + m_impl->events.wait_on_event(best_slot->id, deps); + + spdlog::trace( + "acquired data stream {} for {} (start_time={})", + best_slot->id, + kind, + best_slot->estimated_finish_time + ); + return best_slot->id; +} + +DeviceEvent DeviceDataStreams::release_stream(DeviceStreamId stream_id, uint64_t cost) { + auto it = m_impl->id_to_slot.find(stream_id); + KMM_ASSERT(it != m_impl->id_to_slot.end()); + + auto event = m_impl->events.record(stream_id); + auto& slot = *it->second; + slot.active_users--; + slot.estimated_finish_time += cost; + slot.last_event = event; + m_impl->estimated_finish_time[event] = slot.estimated_finish_time; + + spdlog::trace( + "released data stream {} for {} (finish_time={})", + slot.id, + slot.kind, + slot.estimated_finish_time + ); + + return event; +} + +std::ostream& operator<<(std::ostream& stream, StreamKind kind) { + switch (kind) { + case StreamKind::DeviceToDevice: + return stream << "DeviceToDevice"; + break; + case StreamKind::HostToDevice: + return stream << "HostToDevice"; + break; + case StreamKind::DeviceToHost: + return stream << "DeviceToHost"; + break; + default: + return stream << "Unknown"; + } +} + +} // namespace kmm diff --git a/src/runtime/device_event.cpp b/src/runtime/device_event.cpp new file mode 100644 index 00000000..89f9d858 --- /dev/null +++ b/src/runtime/device_event.cpp @@ -0,0 +1,185 @@ +#include +#include + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/device_event.hpp" +#include "kmm/runtime/device_event_registry.hpp" + +namespace kmm { + +DeviceEventSet::DeviceEventSet(std::initializer_list list) { + *this = list; +} + +DeviceEventSet::DeviceEventSet(const DeviceEvent& event) { + insert(event); +} + +DeviceEventSet& DeviceEventSet::operator=(std::initializer_list list) { + m_events.resize(list.size()); + m_events.clear(); + + for (auto e : list) { + insert(e); + } + + return *this; +} + +void DeviceEventSet::insert(DeviceEvent e) noexcept { + if (e.is_null()) { + return; + } + + bool found = false; + + for (size_t i = 0; i < m_events.size(); i++) { + if (m_events[i].stream() == e.stream()) { + m_events[i] = std::max(m_events[i], e); + found = true; + } + } + + if (KMM_UNLIKELY(!found)) { + KMM_ASSERT(m_events.try_push_back(e)); + } +} + +void DeviceEventSet::insert(const DeviceEventSet& that) noexcept { + size_t n = m_events.size(); + + for (const auto& e : that.m_events) { + bool found = false; + + for (size_t i = 0; i < n; i++) { + if (m_events[i].stream() == e.stream()) { + m_events[i] = std::max(m_events[i], e); + found = true; + } + } + + if (KMM_UNLIKELY(!found)) { + KMM_ASSERT(m_events.try_push_back(e)); + } + } +} + +void DeviceEventSet::insert(DeviceEventSet&& that) noexcept { + if (m_events.is_empty()) { + m_events = std::move(that.m_events); + } else { + insert(that); + } + + that.clear(); +} + +void DeviceEventSet::prune(const DeviceEventRegistry& registry) noexcept { + size_t index = 0; + + while (true) { + if (index >= m_events.size()) { + return; + } + + if (registry.is_ready(m_events[index])) { + break; + } + + index++; + } + + size_t new_size = m_events.size() - 1; + std::swap(m_events[index], m_events[new_size]); + + while (index < new_size) { + if (!registry.is_ready(m_events[index])) { + index++; + } else { + new_size--; + std::swap(m_events[index], m_events[new_size]); + } + } + + m_events.truncate(new_size); +} + +void DeviceEventSet::clear() noexcept { + m_events.clear(); +} + +bool DeviceEventSet::is_empty() const noexcept { + return m_events.is_empty(); +} + +bool DeviceEventSet::contains(const DeviceEvent& event) const noexcept { + if (event.is_null()) { + return true; + } + + for (const auto& e : m_events) { + if (event.precedes(e)) { + return true; + } + } + + return false; +} + +bool DeviceEventSet::contains(const DeviceEventSet& events) const noexcept { + for (const auto& e : events.m_events) { + if (!contains(e)) { + return false; + } + } + + return true; +} + +DeviceEvent DeviceEventSet::find(DeviceStreamId stream_id) const noexcept { + for (const auto& e : m_events) { + if (e.stream() == stream_id) { + return e; + } + } + + return DeviceEvent {}; +} + +std::ostream& operator<<(std::ostream& stream, const DeviceStreamId& e) { + if (e.is_null()) { + return stream << ""; + } else { + return stream << e.get(); + } +} + +std::ostream& operator<<(std::ostream& stream, const DeviceEvent& e) { + if (e.is_null()) { + return stream << ""; + } else { + return stream << e.stream() << ":" << e.index(); + } +} + +std::ostream& operator<<(std::ostream& stream, const DeviceEventSet& e) { + std::vector events = {e.m_events.begin(), e.m_events.end()}; + std::sort(events.begin(), events.end()); + + stream << "{"; + bool is_first = true; + + for (const auto& item : events) { + if (!is_first) { + stream << ", "; + } + + is_first = false; + stream << item; + } + + return stream << "}"; +} + +} // namespace kmm diff --git a/src/runtime/device_event_registry.cpp b/src/runtime/device_event_registry.cpp new file mode 100644 index 00000000..0eb4636b --- /dev/null +++ b/src/runtime/device_event_registry.cpp @@ -0,0 +1,790 @@ +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "ankerl/unordered_dense.h" +#include "fmt/chrono.h" +#include "fmt/ostream.h" +#include "spdlog/spdlog.h" + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/device_event_registry.hpp" +#include "kmm/runtime/device_stream.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +class PrecedenceVector { + public: + void update(DeviceEvent e) { + auto& it = m_entries[e.stream()]; + it = std::max(e.index(), it); + } + + bool contains(DeviceEvent e) const { + auto it = m_entries.find(e.stream()); + return it != m_entries.end() && it->second >= e.index(); + } + + void merge(const PrecedenceVector& that) { + for (auto entry : that.m_entries) { + auto& it = m_entries[entry.first]; + it = std::max(entry.second, it); + } + } + + DeviceEventSet to_set() { + DeviceEventSet result; + result.m_events.resize(m_entries.size()); + size_t i = 0; + + for (auto entry : m_entries) { + result.m_events[i++] = DeviceEvent {entry.first, entry.second}; + } + + return result; + } + + private: + ankerl::unordered_dense::map m_entries; +}; + +struct EventCallback { + EventCallback(uint64_t event_index, NotifyHandle callback) : + event_index(event_index), + callback(std::move(callback)) {} + + friend bool operator<(const EventCallback& a, const EventCallback& b) { + return a.event_index > b.event_index; + } + + uint64_t event_index; + mutable NotifyHandle callback; +}; + +struct StreamState { + KMM_NOT_COPYABLE_OR_MOVABLE(StreamState) + + public: + StreamState(DeviceStreamId id, GPUStreamRef stream_ref, std::string name) : + id(id), + context(stream_ref.context()), + stream(stream_ref.stream()), + stream_key(stream_ref.stream_id()), + name(std::move(name)) {} + + friend std::ostream& operator<<(std::ostream& stream, const StreamState& state) { + stream << state.id; + + if (!state.name.empty()) { + stream << " \"" << state.name << "\""; + } + + return stream; + } + + ~StreamState() { + GPUContextGuard guard {context}; + + for (const auto& entry : pending_events) { + KMM_GPU_CHECK(g_event_synchronize(entry.second)); + KMM_GPU_CHECK(g_event_destroy(entry.second)); + } + + for (g_event_t event : free_events) { + KMM_GPU_CHECK(g_event_destroy(event)); + } + + while (!callbacks.empty()) { + auto callback = std::move(callbacks.top().callback); + callbacks.pop(); + callback.notify_and_clear(); + } + } + + // Precondition: caller holds `mutex`. + g_event_t pop_event_locked() { + if (!free_events.empty()) { + g_event_t event = free_events.back(); + free_events.pop_back(); + return event; + } + + GPUContextGuard guard {context}; + g_event_t event; + KMM_GPU_CHECK(g_event_create(&event, G_EVENT_DISABLE_TIMING)); + return event; + } + + // Precondition: caller holds `mutex`. + g_event_t try_resolve_event_locked(uint64_t event_id) const { + uint64_t first_pending = first_pending_event.load(std::memory_order_acquire); + + if (event_id < first_pending) { + return nullptr; + } + + // Event ids are globally unique (shared across all streams), so this stream's own + // `pending_events` no longer line up with a contiguous range: find the id by binary + // search instead of by direct offset. + auto it = std::lower_bound( + pending_events.begin(), + pending_events.end(), + event_id, + [](const auto& entry, uint64_t id) { return entry.first < id; } + ); + + KMM_ASSERT(it != pending_events.end() && it->first == event_id); + return it->second; + } + + // Precondition: caller holds `mutex`. Drawing the id here (rather than before the lock is + // acquired) guarantees that id order always matches the order events are pushed onto this + // stream's `pending_events`, even if `record()` is ever called concurrently for this stream. + uint64_t record_locked(std::atomic& next_event_id) { + if (released) { + throw std::runtime_error( + "cannot record an event on a stream that has been unregistered" + ); + } + + g_event_t event = pop_event_locked(); + + try { + KMM_GPU_CHECK(g_event_record(event, stream)); + } catch (...) { + // don't leak the event! + free_events.push_back(event); + throw; + } + + uint64_t new_event = next_event_id.fetch_add(1, std::memory_order_relaxed); + spdlog::debug("recorded new event {} on stream {}", DeviceEvent(id, new_event), id); + + // if there are not pending events, we can quickly check if this new event also completed + // immediately. This is possible if the stream is idle, and we just recorded after nothing. + bool is_complete = pending_events.empty() && g_event_query(event) == G_SUCCESS && false; + + if (is_complete) { + free_events.push_back(event); + first_pending_event.store(new_event + 1, std::memory_order_release); + spdlog::debug("completed event {} on stream {}", DeviceEvent(id, new_event), id); + } else { + pending_events.push_back({new_event, event}); + } + + last_pending_event.store(new_event, std::memory_order_release); + + // `preceded_by` (plus this fresh self-entry) now exactly describes the dependencies of + // the event just recorded, so it is back in sync with `last_pending_event`. + preceded_by.update(DeviceEvent {id, new_event}); + preceded_by_is_last = true; + + return new_event; + } + + // Reclaims events that have completed back into the free list and fires callbacks that were + // waiting on them. Precondition: caller holds `state.mutex`. + void make_progress() { + while (!pending_events.empty()) { + auto [completed_id, event] = pending_events.front(); + g_result_t result = g_event_query(event); + + if (result == G_ERROR_NOT_READY) { + break; + } + + KMM_GPU_CHECK(result); + + pending_events.pop_front(); + free_events.push_back(event); + + // A stream's own events always complete in the order they were recorded, so once + // `completed_id` is done, everything this stream recorded before it is done too. + first_pending_event.store(completed_id + 1, std::memory_order_relaxed); + spdlog::debug("completed event {} on stream {}", DeviceEvent(id, completed_id), id); + } + + uint64_t ready_before = first_pending_event.load(std::memory_order_relaxed); + + while (!callbacks.empty() && callbacks.top().event_index < ready_before) { + callbacks.top().callback.notify(); + callbacks.pop(); + } + } + + void trim_event_pool_locked() { + GPUContextGuard context_guard {context}; + + for (g_event_t event : free_events) { + KMM_GPU_CHECK(g_event_synchronize(event)); + KMM_GPU_CHECK(g_event_destroy(event)); + } + + free_events.clear(); + } + + const DeviceStreamId id; + const g_context_t context; + const g_stream_t stream; + + // Canonical driver-level identity of `stream`, cached at registration time so that + // `lookup_stream` can compare against it without a fresh `gpuStreamGetId` call per candidate. + const GPUStreamId stream_key; + + // Optional human-readable name, included in log messages that refer to this stream. + const std::string name; + + mutable std::mutex mutex; + std::vector free_events; + // Pairs of (globally unique event id, g_event_t), ordered by id (equivalently, by the order + // in which they were recorded on this stream). + std::deque> pending_events; + std::priority_queue callbacks; + + std::atomic first_pending_event {1}; + std::atomic last_pending_event {0}; + + // Tracks what this stream has waited for. This allows one to check if some event on another + // stream already precedes this stream's future work. `preceded_by_is_last` indicates if + // the precedence vector corresponds to the last recorded event on this stream or not. + PrecedenceVector preceded_by; + bool preceded_by_is_last = false; + + // Indicates if `unregister_stream` has been called. No new events can be recorded. + bool released = false; +}; + +} // namespace kmm + +template<> +struct fmt::formatter: fmt::ostream_formatter {}; + +namespace kmm { + +struct DeviceEventRegistry::Impl: reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(Impl) + + public: + Impl() = default; + + ~Impl() { + for (auto& slot : streams) { + delete slot.load(std::memory_order_relaxed); + } + } + + StreamState* stream_opt(DeviceStreamId id) { + return id.is_null() ? nullptr : streams[id.get()].load(std::memory_order_acquire); + } + + StreamState& stream(DeviceStreamId id) { + StreamState* state = stream_opt(id); + KMM_ASSERT(state != nullptr); + return *state; + } + + std::array, MAX_DEVICE_STREAMS> streams {}; + + // Shared across all streams so that every recorded event gets a globally unique id, which + // makes events unambiguous when printed/logged (instead of resetting per stream). + std::atomic next_event_id {1}; +}; + +KMM_REFCNT_TRAITS_IMPL(DeviceEventRegistry::Impl) + +DeviceEventRegistry::DeviceEventRegistry() : m_impl(make_refcnt()) {} + +DeviceStreamId DeviceEventRegistry::register_stream( + GPUStreamRef stream_ref, + std::string name +) const { + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + auto id = DeviceStreamId(i); + auto* state = new StreamState(id, stream_ref, name); + StreamState* expected = nullptr; + + if (m_impl->streams[i].compare_exchange_strong( + expected, + state, + std::memory_order_release, + std::memory_order_relaxed + )) { + spdlog::debug("registered new stream {} from {}", *state, state->stream_key); + return id; + } + + delete state; + } + + throw std::runtime_error( + "cannot register stream: the maximum number of streams (MAX_DEVICE_STREAMS) has been reached" + ); +} + +void DeviceEventRegistry::unregister_stream(DeviceStreamId stream_id) const { + StreamState* state = m_impl->streams[stream_id.get()].load(std::memory_order_acquire); + + // no stream registered, just ignore + if (state == nullptr) { + return; + } + + std::lock_guard guard(state->mutex); + + // stream was already released, skip it + if (state->released) { + return; + } + + // Block until every event ever recorded on this stream has completed, so that once this + // returns, nothing is still relying on the stream being alive. Held under `mutex` for the + // whole check-and-act so a concurrent call for the same stream either blocks here and then + // observes `released` and returns, or never gets in until this one has fully finished. + KMM_GPU_CHECK(g_stream_synchronize(state->stream)); + + state->make_progress(); + state->trim_event_pool_locked(); + state->released = true; + + spdlog::debug("unregistered new stream {} from GPU stream {}", state->id, state->stream_key); +} + +std::optional DeviceEventRegistry::lookup_stream(GPUStreamId target) const { + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + StreamState* state = m_impl->streams[i].load(std::memory_order_acquire); + + if (state != nullptr && state->stream_key == target) { + return DeviceStreamId(i); + } + } + + return std::nullopt; +} + +DeviceStreamId DeviceEventRegistry::lookup_or_register_stream(GPUStreamRef stream_ref) const { + GPUStreamId target = stream_ref.stream_id(); + + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + StreamState* state = m_impl->streams[i].load(std::memory_order_acquire); + + if (state == nullptr) { + auto id = DeviceStreamId(i); + auto* new_state = new StreamState(id, stream_ref, ""); + + if (m_impl->streams[i].compare_exchange_strong( + state, + new_state, + std::memory_order_release, + std::memory_order_relaxed + )) { + spdlog::debug( + "registered new stream {} from GPU stream {}", + new_state->id, + new_state->stream_key + ); + return id; + } + + delete new_state; + } else if (state->stream_key == target) { + return DeviceStreamId(i); + } + } + + throw std::runtime_error( + "cannot register stream: the maximum number of streams (MAX_DEVICE_STREAMS) has been reached" + ); +} + +void DeviceEventRegistry::shutdown() const { + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + unregister_stream(DeviceStreamId(i)); + } +} + +g_stream_t DeviceEventRegistry::get(DeviceStreamId stream_id) const { + return m_impl->stream(stream_id).stream; +} + +g_context_t DeviceEventRegistry::context(DeviceStreamId stream_id) const { + return m_impl->stream(stream_id).context; +} + +DeviceStream DeviceEventRegistry::stream(DeviceStreamId stream_id) const { + return DeviceStream(*this, stream_id); +} + +bool DeviceEventRegistry::has_context(DeviceStreamId stream_id, GPUContextId context_id) const { + if (auto* state = m_impl->stream_opt(stream_id)) { + return state->stream_key.context() == context_id; + } else { + return false; + } +} + +DeviceEvent DeviceEventRegistry::record(DeviceStreamId stream_id) const { + StreamState& state = m_impl->stream(stream_id); + std::lock_guard guard(state.mutex); + auto index = state.record_locked(m_impl->next_event_id); + return DeviceEvent {stream_id, index}; +} + +void DeviceEventRegistry::wait_on_event(DeviceStreamId stream_id, DeviceEvent event) const { + if (event.is_null()) { + return; + } + + // Work recorded on a stream is always ordered after events already recorded on that same + // stream, so there is nothing to do here. This also avoids locking this stream's mutex + // against itself below (source and destination would be the same mutex). + if (event.stream() == stream_id) { + return; + } + + StreamState& dst = m_impl->stream(stream_id); + StreamState& src = m_impl->stream(event.stream()); + uint64_t src_event = event.index(); + + std::scoped_lock guard {dst.mutex, src.mutex}; + + // `dst` already waited (directly or transitively) on `event` + if (dst.preceded_by.contains(event)) { + return; + } + + g_event_t handle = src.try_resolve_event_locked(src_event); + + if (handle == nullptr) { + return; + } + + KMM_GPU_CHECK(g_stream_wait_event(dst.stream, handle, 0)); + spdlog::debug("stream {} must wait on event {}", dst.id, event); + + // This stream now directly waits on `src_event`, so it is always a valid predecessor of + // any future work recorded on this stream. This does not record a new event on `dst`, so + // `preceded_by` now describes dependencies for whatever gets recorded next rather than for + // `last_pending_event`, which makes it out of sync until the next `record_locked()` call. + dst.preceded_by.update(event); + dst.preceded_by_is_last = false; + + // We can only fold in everything that precedes `src_event` if `src_event` is `src`'s latest + // event and `src` itself has complete knowledge of what precedes it. + + /* TODO: not sure about this? */ + // uint64_t src_last = src.last_pending_event.load(std::memory_order_relaxed); + // if (src_last == src_event && src.preceded_by_is_last) { + // dst.preceded_by.merge(src.preceded_by); + // } +} + +void DeviceEventRegistry::wait_on_event( + DeviceStreamId stream_id, + const DeviceEventSet& events +) const { + for (const auto& event : events) { + wait_on_event(stream_id, event); + } +} + +void DeviceEventRegistry::wait_on_event(g_stream_t stream, DeviceEvent event) const { + if (event.is_null()) { + return; + } + + StreamState& src = m_impl->stream(event.stream()); + uint64_t src_event = event.index(); + + std::scoped_lock guard {src.mutex}; + g_event_t handle = src.try_resolve_event_locked(src_event); + + if (handle == nullptr) { + return; + } + + KMM_GPU_CHECK(g_stream_wait_event(stream, handle, 0)); +} + +void DeviceEventRegistry::wait_on_event(g_stream_t stream, const DeviceEventSet& events) const { + for (const auto& event : events) { + wait_on_event(stream, event); + } +} + +void DeviceEventRegistry::wait_on_default_stream(DeviceStreamId stream_id) const { + StreamState& dst = m_impl->stream(stream_id); + GPUContextGuard guard {dst.context}; + + g_event_t event; + KMM_GPU_CHECK(g_event_create(&event, G_EVENT_DISABLE_TIMING)); + + try { + // `nullptr` refers to the CUDA legacy default stream in the current context. + KMM_GPU_CHECK(g_event_record(event, nullptr)); + KMM_GPU_CHECK(g_stream_wait_event(dst.stream, event, 0)); + } catch (...) { + KMM_GPU_CHECK(g_event_destroy(event)); + throw; + } + + KMM_GPU_CHECK(g_event_destroy(event)); +} + +bool DeviceEventRegistry::is_ready(DeviceStreamId stream_id) const { + if (stream_id.is_null()) { + return true; + } + + StreamState& state = m_impl->stream(stream_id); + g_result_t result = g_stream_query(state.stream); + + if (result == G_ERROR_NOT_READY) { + return false; + } + + KMM_GPU_CHECK(result); + return true; +} + +bool DeviceEventRegistry::is_ready(DeviceEvent event) const { + if (event.is_null()) { + return true; + } + + StreamState& state = m_impl->stream(event.stream()); + return event.index() < state.first_pending_event.load(std::memory_order_acquire); +} + +bool DeviceEventRegistry::is_ready(const DeviceEventSet& events) const { + for (const auto& e : events) { + if (!is_ready(e)) { + return false; + } + } + + return true; +} + +bool DeviceEventRegistry::is_latest(DeviceEvent event) const { + if (event.is_null()) { + return false; + } + + StreamState& state = m_impl->stream(event.stream()); + return event.index() == state.last_pending_event.load(std::memory_order_acquire); +} + +bool DeviceEventRegistry::is_latest_in(DeviceStreamId stream_id, const DeviceEventSet& deps) const { + if (stream_id.is_null()) { + return false; + } + + for (const auto& event : deps) { + if (event.stream() == stream_id && is_latest(event)) { + return true; + } + } + + return false; +} + +DeviceEvent DeviceEventRegistry::latest_event(DeviceStreamId stream_id) const { + if (stream_id.is_null()) { + return DeviceEvent::null(); + } + + StreamState& state = m_impl->stream(stream_id); + uint64_t index = state.last_pending_event.load(std::memory_order_acquire); + + if (index == 0) { + return DeviceEvent {}; + } + + return DeviceEvent {stream_id, index}; +} + +DeviceEventSet DeviceEventRegistry::snapshot(DeviceStreamId stream_id) const { + if (stream_id.is_null()) { + return {}; + } + + StreamState& state = m_impl->stream(stream_id); + DeviceEventSet result; + + std::lock_guard guard(state.mutex); + return state.preceded_by.to_set(); +} + +void DeviceEventRegistry::synchronize(DeviceStreamId stream_id) const { + if (stream_id.is_null()) { + return; + } + + StreamState& state = m_impl->stream(stream_id); + + auto before = std::chrono::system_clock::now(); + KMM_GPU_CHECK(g_stream_synchronize(state.stream)); + auto after = std::chrono::system_clock::now(); + + auto duration = after - before; + if (duration > std::chrono::milliseconds(1)) { + spdlog::warn("waited for {} to synchronize with stream {}", duration, stream_id); + } + + std::lock_guard guard(state.mutex); + state.make_progress(); +} + +void DeviceEventRegistry::synchronize(DeviceEvent event) const { + if (event.is_null()) { + return; + } + + StreamState& state = m_impl->stream(event.stream()); + g_event_t handle = nullptr; + + { + std::lock_guard guard(state.mutex); + + if (event.index() >= state.first_pending_event.load(std::memory_order_relaxed)) { + handle = state.try_resolve_event_locked(event.index()); + } + } + + if (handle == nullptr) { + return; + } + + auto before = std::chrono::system_clock::now(); + KMM_GPU_CHECK(g_event_synchronize(handle)); + auto after = std::chrono::system_clock::now(); + + auto duration = after - before; + if (duration > std::chrono::milliseconds(1)) { + spdlog::warn("waited for {} to synchronize with event {}", duration, event); + } + + std::lock_guard guard(state.mutex); + state.make_progress(); +} + +void DeviceEventRegistry::synchronize(const DeviceEventSet& events) const { + for (const auto& e : events) { + synchronize(e); + } +} + +void DeviceEventRegistry::synchronize_all() const { + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + StreamState* state = m_impl->streams[i].load(std::memory_order_acquire); + + if (state != nullptr) { + synchronize(DeviceStreamId(i)); + } + } +} + +void DeviceEventRegistry::attach_callback(DeviceEvent event, NotifyHandle callback) const { + if (event.is_null()) { + callback.notify_and_clear(); + return; + } + + StreamState& state = m_impl->stream(event.stream()); + + { + std::lock_guard guard(state.mutex); + + if (event.index() >= state.first_pending_event.load(std::memory_order_relaxed)) { + state.callbacks.emplace(event.index(), std::move(callback)); + return; + } + } + + callback.notify_and_clear(); +} + +void DeviceEventRegistry::make_progress() const { + for (auto& slot : m_impl->streams) { + StreamState* state = slot.load(std::memory_order_acquire); + + if (state == nullptr) { + continue; + } + + std::lock_guard guard(state->mutex); + state->make_progress(); + } +} + +bool DeviceEventRegistry::is_all_ready() const { + for (auto& slot : m_impl->streams) { + StreamState* state = slot.load(std::memory_order_acquire); + + if (state == nullptr) { + continue; + } + + std::lock_guard guard(state->mutex); + + if (!state->pending_events.empty()) { + return false; + } + } + + return true; +} + +void DeviceEventRegistry::trim_event_pool(DeviceStreamId stream_id) const { + if (stream_id.is_null()) { + return; + } + + StreamState& state = m_impl->stream(stream_id); + std::lock_guard guard(state.mutex); + state.trim_event_pool_locked(); +} + +void DeviceEventRegistry::trim_event_pool() const { + for (uint64_t i = 0; i < MAX_DEVICE_STREAMS; i++) { + StreamState* state = m_impl->streams[i].load(std::memory_order_acquire); + + if (state != nullptr) { + trim_event_pool(DeviceStreamId(i)); + } + } +} + +bool DeviceEventRegistry::precedes(const DeviceEvent& a, const DeviceStreamId& b) const { + if (a.is_null() || b.is_null() || a.stream() == b) { + return true; + } + + if (is_ready(a)) { + return true; + } + + StreamState& dst = m_impl->stream(b); + std::lock_guard guard(dst.mutex); + return dst.preceded_by.contains(a); +} + +bool DeviceEventRegistry::precedes(const DeviceEventSet& a, const DeviceStreamId& b) const { + for (const auto& event : a) { + if (!precedes(event, b)) { + return false; + } + } + + return true; +} + +} // namespace kmm diff --git a/src/runtime/device_resources.cpp b/src/runtime/device_resources.cpp deleted file mode 100644 index 53002388..00000000 --- a/src/runtime/device_resources.cpp +++ /dev/null @@ -1,172 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/runtime/device_resources.hpp" - -namespace kmm { - -struct DeviceResources::Device { - KMM_NOT_COPYABLE_OR_MOVABLE(Device) - - public: - Device(GPUContextHandle context) : context(context) {} - - GPUContextHandle context; - std::vector> streams; - std::vector last_used_streams; -}; - -struct DeviceResources::Stream { - KMM_NOT_COPYABLE_OR_MOVABLE(Stream) - - public: - Stream( - DeviceId device_id, - DeviceStream stream, - GPUContextHandle context, - g_stream_t gpu_stream - ) : - context(context), - resource(DeviceInfo(device_id, context), context, gpu_stream), - stream(stream) {} - - GPUContextHandle context; - DeviceResource resource; - DeviceStream stream; - DeviceEvent last_event; -}; - -DeviceResources::DeviceResources( - std::vector contexts, - size_t streams_per_context, - std::shared_ptr stream_manager -) : - m_stream_manager(stream_manager), - m_streams_per_device(streams_per_context) { - KMM_ASSERT(m_streams_per_device > 0); - - for (size_t i = 0; i < contexts.size(); i++) { - m_devices.emplace_back(std::make_unique(contexts[i])); - - for (size_t j = 0; j < m_streams_per_device; j++) { - auto stream = stream_manager->create_stream(contexts[i]); - auto s = std::make_unique( - DeviceId(i), - stream, - contexts[i], - stream_manager->get(stream) - ); - - m_devices[i]->streams.emplace_back(std::move(s)); - m_devices[i]->last_used_streams.push_back(j); - } - } -} - -DeviceResources::~DeviceResources() { - for (const auto& device : m_devices) { - for (const auto& e : device->streams) { - m_stream_manager->wait_until_ready(e->stream); - } - } -} - -size_t DeviceResources::num_contexts() const { - return m_devices.size(); -} - -GPUContextHandle DeviceResources::context(DeviceId device_id) { - KMM_ASSERT(device_id < m_devices.size()); - return m_devices[device_id]->context; -} - -DeviceResources::Stream* DeviceResources::select_stream_for_operation( - DeviceId device_id, - DeviceStreamSet stream_hint, - const DeviceEventSet& deps -) { - static constexpr size_t INVALID = ~size_t(0); - KMM_ASSERT(device_id < m_devices.size()); - auto& device = *m_devices[device_id]; - size_t stream_index = INVALID; - - // Limit available streams to the range 0...streams.size() - stream_hint &= DeviceStreamSet::range(0, device.streams.size()); - - // No stream given, set it to all streams - if (stream_hint.is_empty()) { - stream_hint = DeviceStreamSet::all(); - } - - // Case 1: Find a stream that contains one of the dependencies - if (stream_index == INVALID) { - for (auto i : device.last_used_streams) { - auto e = device.streams[i]->last_event; - - if (stream_hint.contains(i) && std::find(deps.begin(), deps.end(), e) != deps.end()) { - stream_index = i; - break; - } - } - } - - // Case 2: Otherwise, find the last used stream that is contained in `stream_hint` - if (stream_index == INVALID) { - for (auto i : device.last_used_streams) { - if (stream_hint.contains(i)) { - stream_index = i; - break; - } - } - } - - // Case 3: Otherwise, select the last used stream - if (stream_index == INVALID) { - stream_index = device.last_used_streams[0]; - } - - // Push this stream to the back - auto it = std::find( // - device.last_used_streams.begin(), - device.last_used_streams.end(), - stream_index - ); - std::rotate(it, it + 1, device.last_used_streams.end()); - - spdlog::debug("selected stream index {} for operation on GPU {}", stream_index, device_id); - return device.streams[stream_index].get(); -} - -DeviceEvent DeviceResources::submit( - DeviceId device_id, - DeviceStreamSet stream_hint, - DeviceEventSet deps, - DeviceResourceOperation& op, - std::vector accessors -) { - auto& state = *select_stream_for_operation(device_id, stream_hint, deps); - - try { - GPUContextGuard guard {state.context}; - m_stream_manager->wait_for_events(state.stream, deps); - - op.execute(state.resource, std::move(accessors)); - - m_stream_manager->wait_on_default_stream(state.stream); - auto event = m_stream_manager->record_event(state.stream); - - state.last_event = event; - return event; - } catch (const std::exception& e) { - try { - m_stream_manager->wait_until_ready(state.stream); - } catch (...) { - KMM_PANIC_FMT("fatal error: {}", e.what()); - } - - throw; - } -} - -} // namespace kmm diff --git a/src/runtime/identifiers.cpp b/src/runtime/identifiers.cpp new file mode 100644 index 00000000..7310d121 --- /dev/null +++ b/src/runtime/identifiers.cpp @@ -0,0 +1,67 @@ +#include +#include +#include + +#include "kmm/runtime/identifiers.hpp" + +namespace kmm { + +static std::string to_lower(const std::string& input) { + std::string result = input; + std::transform(result.begin(), result.end(), result.begin(), [](unsigned char c) { + return std::tolower(c); + }); + return result; +} + +static MemoryId parse_memory_id(const std::string& name) { + std::string lower = to_lower(name); + + if (lower == "host" || lower == "cpu") { + return MemoryId::host(); + } + + size_t sep = lower.find(':'); + std::string prefix = sep == std::string::npos ? lower : lower.substr(0, sep); + + if (prefix == "gpu" || prefix == "cuda" || prefix == "hip" || prefix == "device") { + if (sep == std::string::npos) { + return MemoryId::device(DeviceId(0)); + } + + std::string suffix = lower.substr(sep + 1); + + try { + size_t pos; + size_t device_index = std::stoul(suffix, &pos); + + if (pos == suffix.size()) { + return MemoryId::device(DeviceId(device_index)); + } + } catch (const std::exception&) { + // fallthrough to error below + } + } + + throw std::runtime_error("invalid memory identifier: " + name); +} + +MemoryId::MemoryId(const std::string& name) : MemoryId(parse_memory_id(name)) {} + +std::ostream& operator<<(std::ostream& stream, const DeviceId& e) { + return stream << "GPU(" << e.get() << ")"; +} + +std::ostream& operator<<(std::ostream& stream, const BufferId& e) { + return stream << "Buffer(" << e.get() << ")"; +} + +std::ostream& operator<<(std::ostream& stream, const MemoryId& e) { + if (e.is_host()) { + return stream << "Host"; + } else { + return stream << e.as_device(); + } +} + +} // namespace kmm diff --git a/src/runtime/memops/copy.cpp b/src/runtime/memops/copy.cpp new file mode 100644 index 00000000..02549bd0 --- /dev/null +++ b/src/runtime/memops/copy.cpp @@ -0,0 +1,151 @@ +#include +#include + +#include "simplify_dims.hpp" + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/copy.hpp" + +namespace kmm { + +static Range dim_offset_range( + ptrdiff_t base_offset, + const CopyDim* dims, + size_t num_dims, + size_t element_size, + memops_stride_type CopyDim::* stride_member +) { + ptrdiff_t lo = base_offset; + ptrdiff_t hi = base_offset; + + for (size_t i = 0; i < num_dims; i++) { + if (dims[i].extent < 1) { + return {base_offset, base_offset}; + } + + ptrdiff_t span = checked_mul(dims[i].extent - 1, dims[i].*stride_member); + lo = checked_add(lo, span < 0 ? span : 0); + hi = checked_add(hi, span > 0 ? span : 0); + } + + return {lo, checked_add(hi, element_size)}; +} + +Range CopyDescription::src_range() const { + return dim_offset_range( + static_cast(src_offset), + dims, + num_dims, + element_size, + &CopyDim::src_stride + ); +} + +Range CopyDescription::dst_range() const { + return dim_offset_range( + static_cast(dst_offset), + dims, + num_dims, + element_size, + &CopyDim::dst_stride + ); +} + +CopyDescription CopyDescription::simplify() const { + CopyDescription result = *this; + + for (size_t i = 0; i < result.num_dims; i++) { + CopyDim& dim = result.dims[i]; + + // zero extent means no copy at all. + if (dim.extent <= 0) { + return CopyDescription {0}; + } + + // Rewrite a negative primary stride to positive + if (dim.dst_stride < 0) { + result.src_offset += (dim.extent - 1) * dim.src_stride; + result.dst_offset += (dim.extent - 1) * dim.dst_stride; + + dim.src_stride = -dim.src_stride; + dim.dst_stride = -dim.dst_stride; + } + } + + result.num_dims = simplify_dims( + result.dims, + result.num_dims, + result.dims, + [](const CopyDim& a, const CopyDim& b) { + // Descending stride order, keyed on `dst_stride` (ties broken by `src_stride`): it + // matters more that writes to the innermost axis coalesce than that reads do. + return a.dst_stride != b.dst_stride + ? unsigned_abs(a.dst_stride) > unsigned_abs(b.dst_stride) + : unsigned_abs(a.src_stride) > unsigned_abs(b.src_stride); + }, + [](CopyDim& outer, const CopyDim& inner) { + if (outer.src_stride == inner.src_stride * inner.extent + && outer.dst_stride == inner.dst_stride * inner.extent) { + outer.extent *= inner.extent; + outer.src_stride = inner.src_stride; + outer.dst_stride = inner.dst_stride; + return true; + } + + return false; + } + ); + + // The innermost axis (last, since `dims` is sorted in descending stride order) may itself be + // contiguous with the element: if its stride on both sides equals `element_size`, it can be + // folded into `element_size` instead of being kept as a separate axis. + while (result.num_dims > 0 + && is_equal(result.dims[result.num_dims - 1].src_stride, result.element_size) + && is_equal(result.dims[result.num_dims - 1].dst_stride, result.element_size)) { + result.element_size *= checked_cast(result.dims[result.num_dims - 1].extent); + result.num_dims--; + } + + return result; +} + +namespace memops { + +static void copy_dim( + const std::byte* src, + std::byte* dst, + const CopyDim* dims, + size_t num_dims, + size_t element_size +) { + if (num_dims == 0) { + std::memcpy(dst, src, element_size); + return; + } + + for (memops_extent_type i = 0; i < dims->extent; i++) { + copy_dim( + src + i * dims->src_stride, + dst + i * dims->dst_stride, + dims + 1, + num_dims - 1, + element_size + ); + } +} + +void copy(const void* src_addr, void* dst_addr, const CopyDescription& description) { + copy_dim( + static_cast(src_addr) + description.src_offset, + static_cast(dst_addr) + description.dst_offset, + description.dims, + description.num_dims, + description.element_size + ); +} + +} // namespace memops + +} // namespace kmm diff --git a/src/runtime/memops/copy_gpu.cu b/src/runtime/memops/copy_gpu.cu new file mode 100644 index 00000000..181b1864 --- /dev/null +++ b/src/runtime/memops/copy_gpu.cu @@ -0,0 +1,489 @@ +#include +#include +#include +#include +#include +#include + +#include "memops_gpu_kernels.cuh" + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/vec.hpp" +#include "kmm/runtime/memops/copy_gpu.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm::memops { + +struct CopyPlan { + const void* src_addr; + void* dst_addr; + size_t line_width; + size_t num_dims; + size_t extents[MEMOPS_MAX_DIMS + 1] = {}; + ptrdiff_t input_strides[MEMOPS_MAX_DIMS + 1] = {}; + ptrdiff_t output_strides[MEMOPS_MAX_DIMS + 1] = {}; + + template + bool is_address_aligned() const { + if (!is_divisible(reinterpret_cast(src_addr), alignof(T))) { + return false; + } + + if (!is_divisible(reinterpret_cast(dst_addr), alignof(T))) { + return false; + } + + for (size_t i = 0; i < num_dims; i++) { + if (!is_divisible(input_strides[i], alignof(T))) { + return false; + } + + if (!is_divisible(output_strides[i], alignof(T))) { + return false; + } + } + + return true; + } + + template + bool is_aligned() const { + if (!is_divisible(line_width, sizeof(T))) { + return false; + } + + return is_address_aligned(); + } +}; + +CopyPlan make_plan(const void* src_base, void* dst_base, const CopyDescription& description) { + CopyPlan plan; + plan.num_dims = 0; + plan.src_addr = static_cast(src_base) + description.src_offset; + plan.dst_addr = static_cast(dst_base) + description.dst_offset; + plan.line_width = description.element_size; + + size_t old_rank = description.num_dims; + size_t new_rank = 0; + CopyDim dims[MEMOPS_MAX_DIMS] = {}; + std::copy_n(description.dims, old_rank, dims); + + for (size_t i = 0; i < old_rank; i++) { + for (size_t j = i + 1; j < old_rank; j++) { + if (unsigned_abs(dims[j].dst_stride) < unsigned_abs(dims[i].dst_stride)) { + std::swap(dims[i], dims[j]); + } + } + + auto new_dim = dims[i]; + + // no copy needed at all, just set line width to zero bytes. + if (new_dim.extent <= 0) { + plan.line_width = 0; + new_rank = 0; + break; + } + + // skip this dimension if it has extent of one + if (new_dim.extent == 1) { + continue; + } + + // if the dst_stride is zero, all values land at the same location. We can effectively consider its + // extent to be equal to one. + if (new_dim.dst_stride == 0) { + // TODO: Maybe throw an exception here? Why would you want a dst stride of zero? + continue; + } + + // fix negative stride by subtracting offset from the pointer + if (new_dim.dst_stride < 0) { + plan.src_addr = static_cast(plan.src_addr) + + (new_dim.extent - 1) * new_dim.src_stride; + plan.dst_addr = + static_cast(plan.dst_addr) + (new_dim.extent - 1) * new_dim.dst_stride; + + new_dim.dst_stride = -new_dim.dst_stride; + new_dim.src_stride = -new_dim.src_stride; + } + + if (new_rank == 0) { + // if the line width equals the stride, we can just extend the line width. + if (is_equal(plan.line_width, new_dim.src_stride) + && is_equal(plan.line_width, new_dim.dst_stride)) { + plan.line_width *= checked_cast(new_dim.extent); + continue; + } + } else { + auto k = new_rank - 1; + + if (is_equal(new_dim.dst_stride, plan.output_strides[k] * plan.extents[k]) + && is_equal(new_dim.src_stride, plan.input_strides[k] * plan.extents[k])) { + plan.extents[new_rank - 1] *= new_dim.extent; + continue; + } + + if (is_divisible(new_dim.dst_stride, plan.output_strides[k]) + && is_divisible(new_dim.src_stride, plan.input_strides[k])) { + // TODO: should be something smart when the strides are multiples of each other + } + } + + plan.extents[new_rank] = new_dim.extent; + plan.input_strides[new_rank] = new_dim.src_stride; + plan.output_strides[new_rank] = new_dim.dst_stride; + new_rank++; + } + + plan.num_dims = new_rank; + return plan; +} + +template +void launch_strided_kernel_rank_typed_recur( + g_stream_t stream, + Vec extents, + std::byte* dst_addr, + const std::byte* src_addr, + Vec dst_strides, + Vec src_strides +) { + constexpr uint32_t max_blocks = 1024; + constexpr uint32_t threads_per_block = 256; + + // stride 0 must be contiguous + KMM_ASSERT(dst_strides[0] == sizeof(T)); + KMM_ASSERT(src_strides[0] == sizeof(T)); + + Vec local_extents; + uint32_t num_threads = 1; + + for (size_t i = 0; i < N; i++) { + if (extents[i] == 0) { + return; + } + + // Largest count along axis `i` that keeps the running thread total within the range + // that `FastDivisor` accepts as a numerator. + uint32_t chunk = IndexMapper::max_volume / num_threads; + + while (extents[i] > chunk) { + Vec head = extents; + head[i] = chunk; + + launch_strided_kernel_rank_typed_recur( + stream, + head, + dst_addr, + src_addr, + dst_strides, + src_strides + ); + + dst_addr += ptrdiff_t(chunk) * dst_strides[i]; + src_addr += ptrdiff_t(chunk) * src_strides[i]; + extents[i] -= chunk; + } + + local_extents[i] = static_cast(extents[i]); // safe cast + num_threads *= static_cast(extents[i]); + } + + uint32_t grid_size = std::min(div_ceil(num_threads, threads_per_block), max_blocks); + auto mapper = IndexMapper(local_extents); + + elementwise_copy_kernel<<>>( + dst_addr, + src_addr, + mapper, + dst_strides, + src_strides + ); +} + +template +void launch_transpose_kernel_rank_typed_recur( + g_stream_t stream, + Vec extents, + std::byte* dst_addr, + const std::byte* src_addr, + Vec dst_strides, + Vec src_strides +) { + static_assert(N >= 2, "transpose kernel must have at least 2 dimensions"); + constexpr uint32_t max_blocks = 1024; + constexpr uint32_t tile_size = 32; + constexpr uint32_t block_dim_x = 32; + constexpr uint32_t block_dim_y = 8; + + Vec grid_extents; + uint32_t num_blocks = 1; + + for (size_t i = 0; i < 2; i++) { + if (extents[i] == 0) { + return; + } + + size_t count = div_ceil(extents[i], size_t(tile_size)); + uint32_t chunk = IndexMapper::max_volume / num_blocks; + + chunk = std::min(chunk, IndexMapper::max_volume / tile_size); + + while (count > chunk) { + Vec head = extents; + head[i] = size_t(chunk) * tile_size; + + launch_transpose_kernel_rank_typed_recur( + stream, + head, + dst_addr, + src_addr, + dst_strides, + src_strides + ); + + dst_addr += ptrdiff_t(head[i]) * dst_strides[i]; + src_addr += ptrdiff_t(head[i]) * src_strides[i]; + extents[i] -= head[i]; + count -= chunk; + } + + grid_extents[i] = static_cast(count); + num_blocks *= static_cast(count); + } + + for (size_t i = 2; i < N; i++) { + if (extents[i] == 0) { + return; + } + + size_t count = extents[i]; + uint32_t chunk = IndexMapper::max_volume / num_blocks; + + while (count > chunk) { + Vec head = extents; + head[i] = size_t(chunk); + + launch_transpose_kernel_rank_typed_recur( + stream, + head, + dst_addr, + src_addr, + dst_strides, + src_strides + ); + + dst_addr += ptrdiff_t(head[i]) * dst_strides[i]; + src_addr += ptrdiff_t(head[i]) * src_strides[i]; + extents[i] -= head[i]; + count -= chunk; + } + + grid_extents[i] = static_cast(count); + num_blocks *= static_cast(count); + } + + uint32_t grid_size = std::min(num_blocks, max_blocks); + auto mapper = IndexMapper(grid_extents); + + transpose_copy_kernel + <<>>( + dst_addr, + src_addr, + static_cast(extents[0]), + static_cast(extents[1]), + mapper, + dst_strides, + src_strides + ); +} + +template +void launch_strided_kernel_rank_typed(g_stream_t stream, const CopyPlan& plan) { + KMM_ASSERT(plan.num_dims == N); + KMM_ASSERT(plan.line_width % sizeof(T) == 0); + + Vec extents; + Vec src_strides; + Vec dst_strides; + + extents[0] = plan.line_width / sizeof(T); + src_strides[0] = sizeof(T); + dst_strides[0] = sizeof(T); + + for (size_t i = 0; i < N; i++) { + extents[i + 1] = plan.extents[i]; + src_strides[i + 1] = plan.input_strides[i]; + dst_strides[i + 1] = plan.output_strides[i]; + } + + auto* dst_addr = static_cast(plan.dst_addr); + const auto* src_addr = static_cast(plan.src_addr); + + if constexpr (N >= 2) { + // if N >= 2, we can check if this copy is actually a transposition. To detect this, we scan over the + // strides and attempt to find the "unit" axes for the source and destination. This is the axis that meets + // the following criteria: + // - Must be sufficiently large + // - The stride cannot be zero + // - The stride is contiguous (or very close to contiguous). + // + // If the unit source axis is different from the unit destination axis, then we have a transposition and we + // call the special transposition kernel to handle this. + const size_t minimum_length = 32; + const size_t near_contiguous = 4 * sizeof(T); + + size_t src_unit = 0; + size_t dst_unit = 0; + + for (size_t i = 1; i < N + 1; i++) { + if (extents[i] >= minimum_length && src_strides[i] != 0 + && unsigned_abs(src_strides[i]) <= near_contiguous + && (src_unit == 0 + || unsigned_abs(src_strides[i]) < unsigned_abs(src_strides[src_unit]))) { + src_unit = i; + } + if (extents[i] >= minimum_length && dst_strides[i] != 0 + && unsigned_abs(dst_strides[i]) <= near_contiguous + && (dst_unit == 0 + || unsigned_abs(dst_strides[i]) < unsigned_abs(dst_strides[dst_unit]))) { + dst_unit = i; + } + } + + if (src_unit != 0 && dst_unit != 0 && src_unit != dst_unit) { + // A genuine transpose never merges anything into the line-width axis, so it stays a + // single element. In that case drop it and dispatch over the N real axes, so e.g. a + // 2D transpose runs the rank-2 kernel rather than a rank-3 one with a trailing 1. + const bool drop_line_axis = extents[0] == 1; + + // rotate the src-contiguous axis to position 0 + for (size_t i = src_unit; i > 0; i--) { + std::swap(src_strides[i], src_strides[i - 1]); + std::swap(dst_strides[i], dst_strides[i - 1]); + std::swap(extents[i], extents[i - 1]); + } + + // that rotation shifted every axis below src_unit up by one + if (dst_unit < src_unit) { + dst_unit += 1; + } + + // rotate the dst-contiguous axis to position 1 + for (size_t i = dst_unit; i > 1; i--) { + std::swap(src_strides[i], src_strides[i - 1]); + std::swap(dst_strides[i], dst_strides[i - 1]); + std::swap(extents[i], extents[i - 1]); + } + + launch_transpose_kernel_rank_typed_recur( + stream, + extents, + dst_addr, + src_addr, + dst_strides, + src_strides + ); + + return; + } + } + + launch_strided_kernel_rank_typed_recur( + stream, + extents, + dst_addr, + src_addr, + dst_strides, + src_strides + ); +} + +template +void launch_strided_kernel_typed(g_stream_t stream, const CopyPlan& plan) { + size_t num_dims = plan.num_dims; + + if (num_dims == 1) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 2) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 3) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 4) { + launch_strided_kernel_rank_typed(stream, plan); + } else { + // should not happen + KMM_PANIC("invalid dimensionality"); + } +} + +bool launch_strided_kernel(g_stream_t stream, const CopyPlan& plan) { + if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else { + KMM_ASSERT(plan.is_aligned()); + launch_strided_kernel_typed(stream, plan); + } + + return true; +} + +void execute_copy_plan(g_stream_t stream, const CopyPlan& plan) { + // nothing to do + if (plan.line_width == 0) { + return; + } + + // simple 1D copy + if (plan.num_dims == 0) { + KMM_GPU_CHECK(g_memcpy_async( + reinterpret_cast(plan.dst_addr), + reinterpret_cast(const_cast(plan.src_addr)), + plan.line_width, + stream + )); + + return; + } + + // 2D copy (if possible) + if (plan.num_dims == 1 && is_greater(plan.input_strides[0], plan.line_width) + && is_greater(plan.output_strides[0], plan.line_width)) { + gpu_memcpy2d_t p; + ::memset(&p, 0, sizeof(gpu_memcpy2d_t)); + + p.srcMemoryType = G_MEMORYTYPE_DEVICE; + p.srcDevice = reinterpret_cast(const_cast(plan.src_addr)); + p.srcPitch = checked_cast(plan.input_strides[0]); + p.dstMemoryType = G_MEMORYTYPE_DEVICE; + p.dstDevice = reinterpret_cast(plan.dst_addr); + p.dstPitch = checked_cast(plan.output_strides[0]); + p.WidthInBytes = checked_cast(plan.line_width); + p.Height = checked_cast(plan.extents[0]); + + KMM_GPU_CHECK(g_memcpy_2d_async(&p, stream)); + return; + } + + launch_strided_kernel(stream, plan); +} + +void copy_gpu( + g_stream_t stream, + const void* src_base, + void* dst_base, + const CopyDescription& description +) { + auto plan = make_plan(src_base, dst_base, description); + execute_copy_plan(stream, plan); +} + +} // namespace kmm::memops diff --git a/src/runtime/memops/fill.cpp b/src/runtime/memops/fill.cpp new file mode 100644 index 00000000..1366867b --- /dev/null +++ b/src/runtime/memops/fill.cpp @@ -0,0 +1,116 @@ +#include +#include + +#include "simplify_dims.hpp" + +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/fill.hpp" + +namespace kmm { + +void FillDescription::add_dimension(memops_extent_type extent, memops_stride_type stride) { + // add the dimensions + if (num_dims < MEMOPS_MAX_DIMS) { + dims[num_dims] = FillDim {extent, stride}; + num_dims++; + return; + } + + // could not add the dimensions, try to fuse it with one of the existing dimensions + for (auto& dim : this->dims) { + // `dim` is the outer neighbor of the new axis or has extent 1 + if (dim.stride == stride * extent || dim.extent == 1) { + dim.extent *= extent; + dim.stride = stride; + return; + } + + // `dim` is the inner neighbor of the new axis + if (stride == dim.stride * dim.extent) { + dim.extent *= extent; + return; + } + } + + throw std::runtime_error( + "cannot add dimension to `FillDescription`, exceeds maximum number of dimensions" + ); +} + +FillDescription FillDescription::simplify() const { + FillDescription result = *this; + + for (size_t i = 0; i < result.num_dims; i++) { + FillDim& dim = result.dims[i]; + + // A zero or negative extent makes the whole fill empty. `simplify_dims` handles this. + if (dim.extent <= 0) { + continue; + } + + // Rewrite negative strides to positive, shifting `offset` to the far end of the axis so + // the same elements are still visited. (`simplify_dims` handles negative extents.) + if (dim.stride < 0) { + result.offset += (dim.extent - 1) * dim.stride; + dim.stride = -dim.stride; + } + + // A zero stride revisits the same address `extent` times; for `fill` those repeated + // writes are redundant, so collapse the axis to a single element and let `simplify_dims` + // drop it. + if (dim.stride == 0) { + dim.extent = 1; + } + } + + result.num_dims = simplify_dims( + result.dims, + result.num_dims, + result.dims, + [](const FillDim& a, const FillDim& b) { return a.stride > b.stride; }, + [](FillDim& outer, const FillDim& inner) { + if (outer.stride == inner.stride * inner.extent) { + outer.extent *= inner.extent; + outer.stride = inner.stride; + return true; + } + + return false; + } + ); + + return result; +} + +namespace memops { + +static void fill_dim( + std::byte* dst, + const FillDim* dims, + size_t num_dims, + size_t element_size, + const void* fill_value +) { + if (num_dims == 0) { + std::memcpy(dst, fill_value, element_size); + return; + } + + for (memops_extent_type i = 0; i < dims->extent; i++) { + fill_dim(dst + i * dims->stride, dims + 1, num_dims - 1, element_size, fill_value); + } +} + +void fill(void* dst_addr, const FillDescription& description) { + fill_dim( + static_cast(dst_addr) + description.offset, + description.dims, + description.num_dims, + description.value.length, + description.value.buffer + ); +} + +} // namespace memops + +} // namespace kmm diff --git a/src/runtime/memops/fill_gpu.cu b/src/runtime/memops/fill_gpu.cu new file mode 100644 index 00000000..4f9b8537 --- /dev/null +++ b/src/runtime/memops/fill_gpu.cu @@ -0,0 +1,391 @@ +#include +#include +#include +#include +#include +#include + +#include "memops_gpu_kernels.cuh" + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/point.hpp" +#include "kmm/core/vec.hpp" +#include "kmm/runtime/memops/fill_gpu.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm::memops { + +struct FillPlan { + void* dst_addr; + FillValue fill_pattern; + size_t line_width; + size_t num_dims; + size_t extents[MEMOPS_MAX_DIMS + 1] = {}; + ptrdiff_t strides[MEMOPS_MAX_DIMS + 1] = {}; + + template + T pattern_as() const { + KMM_ASSERT(line_width % sizeof(T) == 0); + KMM_ASSERT(sizeof(T) % fill_pattern.length == 0); + + std::byte buffer[sizeof(T)]; + + for (size_t i = 0; i < sizeof(T); i++) { + buffer[i] = fill_pattern.buffer[i % fill_pattern.length]; + } + + T value; + ::memcpy(&value, buffer, sizeof(T)); + return value; + } + + template + bool is_aligned() const { + // A T-sized store must cover a whole number of pattern periods. + if (!is_divisible(sizeof(T), fill_pattern.length)) { + return false; + } + + if (!is_divisible(line_width, sizeof(T))) { + return false; + } + + if (!is_divisible(reinterpret_cast(dst_addr), alignof(T))) { + return false; + } + + for (size_t i = 0; i < num_dims; i++) { + if (!is_divisible(strides[i], alignof(T))) { + return false; + } + } + + return true; + } +}; + +FillPlan make_plan(void* dst_base, const FillDescription& description) { + FillPlan plan; + plan.num_dims = 0; + plan.dst_addr = static_cast(dst_base) + description.offset; + plan.fill_pattern = description.value; + plan.line_width = description.value.length; + + for (size_t k : std::array {16, 8, 4, 2, 1}) { + bool is_periodic = plan.fill_pattern.length % k == 0; + + for (size_t i = k; i < plan.fill_pattern.length; i++) { + is_periodic &= plan.fill_pattern.buffer[i] == plan.fill_pattern.buffer[i - k]; + } + + if (is_periodic) { + plan.fill_pattern.length = k; + } + } + + size_t old_rank = description.num_dims; + size_t new_rank = 0; + FillDim dims[MEMOPS_MAX_DIMS] = {}; + std::copy_n(description.dims, old_rank, dims); + + for (size_t i = 0; i < old_rank; i++) { + for (size_t j = i + 1; j < old_rank; j++) { + if (unsigned_abs(dims[j].stride) < unsigned_abs(dims[i].stride)) { + std::swap(dims[i], dims[j]); + } + } + + auto new_dim = dims[i]; + + // no copy needed at all, just set line width to zero bytes. + if (new_dim.extent <= 0) { + plan.line_width = 0; + new_rank = 0; + break; + } + + // skip this dimension + if (new_dim.extent == 1) { + continue; + } + + // fix negative stride by subtracting offset from the pointer + if (new_dim.stride < 0) { + plan.dst_addr = + static_cast(plan.dst_addr) + (new_dim.extent - 1) * new_dim.stride; + new_dim.stride = -new_dim.stride; + } + + if (new_rank == 0) { + // if the line width equals the stride, we can just extent the line width. + if (is_equal(plan.line_width, new_dim.stride)) { + plan.line_width *= checked_cast(new_dim.extent); + continue; + } + } else { + if (is_equal(new_dim.stride, plan.strides[new_rank - 1] * plan.extents[new_rank - 1])) { + plan.extents[new_rank - 1] *= new_dim.extent; + continue; + } + } + + plan.extents[new_rank] = new_dim.extent; + plan.strides[new_rank] = new_dim.stride; + new_rank++; + } + + plan.num_dims = new_rank; + return plan; +} + +// Launch `elementwise_fill_kernel` over an index space whose axis 0 is the contiguous run and +// whose remaining axes are strided. The space is peeled apart along each axis so that no single +// launch has more threads than `IndexMapper` can invert (mirrors `ParallelFor::launch_recur`). +template +void launch_strided_fill_recur( + g_stream_t stream, + std::byte* dst_addr, + T value, + Vec extents, + Vec strides +) { + constexpr uint32_t max_blocks = 1024; + constexpr uint32_t threads_per_block = 256; + Vec local_extents; + uint32_t num_threads = 1; + + for (size_t i = 0; i < Rank; i++) { + if (extents[i] == 0) { + return; + } + + // Largest count along axis `i` that keeps the running thread total within the range + // that `FastDivisor` accepts as a numerator. + uint32_t chunk = IndexMapper::max_volume / num_threads; + + while (extents[i] > chunk) { + Vec head = extents; + head[i] = chunk; + launch_strided_fill_recur(stream, dst_addr, value, head, strides); + + dst_addr += ptrdiff_t(chunk) * strides[i]; + extents[i] -= chunk; + } + + local_extents[i] = static_cast(extents[i]); // safe cast + num_threads *= local_extents[i]; + } + + uint32_t grid_size = std::min(div_ceil(num_threads, threads_per_block), max_blocks); + auto mapper = IndexMapper(local_extents); + + KMM_ASSERT(Rank > 0 && strides[0] == sizeof(T)); + elementwise_fill_kernel + <<>>(dst_addr, value, mapper, strides); +} + +template +void launch_strided_kernel_rank_typed(g_stream_t stream, const FillPlan& plan) { + KMM_ASSERT(plan.num_dims + 1 == Rank); + Vec extents; + Vec strides; + + extents[0] = plan.line_width / sizeof(T); + strides[0] = ptrdiff_t(sizeof(T)); + + for (size_t i = 1; i < Rank; i++) { + extents[i] = plan.extents[i - 1]; + strides[i] = plan.strides[i - 1]; + } + + launch_strided_fill_recur( + stream, + static_cast(plan.dst_addr), + plan.pattern_as(), + extents, + strides + ); +} + +template +void launch_strided_kernel_typed(g_stream_t stream, const FillPlan& plan) { + size_t num_dims = plan.num_dims; + + if (num_dims == 0) { + // a 0D fill is rare since it will typically (always?) be performed by cuMemsetD32Async + // instead, we turn it into a 1D fill since that will reduce the number of kernels + // that are precompiled. + FillPlan p = plan; + p.extents[0] = 1; + p.strides[0] = 0; + p.num_dims++; + launch_strided_kernel_typed(stream, p); + } else if (num_dims == 1) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 2) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 3) { + launch_strided_kernel_rank_typed(stream, plan); + } else if (num_dims == 4) { + launch_strided_kernel_rank_typed(stream, plan); + } else { + // should not happen + KMM_PANIC("invalid dimensionality"); + } +} + +bool try_launch_strided_kernel(g_stream_t stream, const FillPlan& plan) { + if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else if (plan.is_aligned()) { + launch_strided_kernel_typed(stream, plan); + } else { + return false; + } + + return true; +} + +bool try_fill_plan(g_stream_t stream, const FillPlan& plan) { + if (plan.line_width == 0) { + return true; + } + + if (plan.num_dims == 0) { + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d32_async( + reinterpret_cast(plan.dst_addr), + plan.pattern_as(), + plan.line_width / sizeof(uint32_t), + stream + )); + + return true; + } + + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d16_async( + reinterpret_cast(plan.dst_addr), + plan.pattern_as(), + plan.line_width / sizeof(uint16_t), + stream + )); + + return true; + } + + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d8_async( + reinterpret_cast(plan.dst_addr), + plan.pattern_as(), + plan.line_width, + stream + )); + + return true; + } + } else if (plan.num_dims == 1 && is_greater(plan.strides[0], plan.line_width)) { + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d2d32_async( + reinterpret_cast(plan.dst_addr), + checked_cast(plan.strides[0]), + plan.pattern_as(), + plan.line_width / sizeof(uint32_t), + checked_cast(plan.extents[0]), + stream + )); + + return true; + } + + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d2d16_async( + reinterpret_cast(plan.dst_addr), + checked_cast(plan.strides[0]), + plan.pattern_as(), + plan.line_width / sizeof(uint16_t), + checked_cast(plan.extents[0]), + stream + )); + + return true; + } + + if (plan.is_aligned()) { + KMM_GPU_CHECK(g_memset_d2d8_async( + reinterpret_cast(plan.dst_addr), + checked_cast(plan.strides[0]), + plan.pattern_as(), + plan.line_width, + checked_cast(plan.extents[0]), + stream + )); + + return true; + } + } + + return try_launch_strided_kernel(stream, plan); +} + +// Fill a pattern that no single element type can handle: its length is not a power of two +void fill_split_pattern(g_stream_t stream, const FillPlan& plan) { + size_t pattern_length = plan.fill_pattern.length; + KMM_ASSERT(pattern_length >= 2); + + size_t offset = 0; + size_t num_elements = plan.line_width / pattern_length; + + if (num_elements != 1 && plan.num_dims == MEMOPS_MAX_DIMS) { + throw std::runtime_error("fill plan too complex"); + } + + while (offset < pattern_length) { + size_t chunk = 32; + + while (true) { + if (chunk <= pattern_length - offset) { + FillPlan sub = plan; + sub.dst_addr = static_cast(plan.dst_addr) + offset; + sub.line_width = chunk; + sub.fill_pattern.length = chunk; + std::copy_n(plan.fill_pattern.buffer + offset, chunk, sub.fill_pattern.buffer); + + if (num_elements != 1) { + sub.num_dims = plan.num_dims + 1; + sub.extents[0] = num_elements; + sub.strides[0] = checked_cast(pattern_length); + std::copy_n(plan.extents, plan.num_dims, sub.extents + 1); + std::copy_n(plan.strides, plan.num_dims, sub.strides + 1); + } + + if (try_fill_plan(stream, sub)) { + offset += chunk; + break; + } + } + + chunk /= 2; + } + } +} + +void fill_gpu(g_stream_t stream, void* dst_base, const FillDescription& description) { + auto plan = make_plan(dst_base, description); + + if (try_fill_plan(stream, plan)) { + return; + } + + fill_split_pattern(stream, plan); +} + +} // namespace kmm::memops diff --git a/src/runtime/memops/memops_gpu_kernels.cuh b/src/runtime/memops/memops_gpu_kernels.cuh new file mode 100644 index 00000000..eab71e1e --- /dev/null +++ b/src/runtime/memops/memops_gpu_kernels.cuh @@ -0,0 +1,351 @@ +#pragma once + +#include +#include +#include + +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/vec.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm::memops { + +//=========================================================================== +// Copy kernels (see copy_gpu.cu) +//=========================================================================== + +template +__global__ void transpose_copy_kernel( + std::byte* dst_addr, + const std::byte* src_addr, + uint32_t extent0, + uint32_t extent1, + IndexMapper mapper, + Vec dst_strides, + Vec src_strides +) { + static_assert(TileSize % BlockSizeX == 0, "invalid block size x"); + static_assert(TileSize % BlockSizeY == 0, "invalid block size y"); + + __shared__ T shared_tile[TileSize][TileSize + 1]; + uint32_t linear_index = blockIdx.x; + Vec p; + + while (mapper.unravel(linear_index, p)) { + std::byte* dst = dst_addr; + const std::byte* src = src_addr; + + // skip axis 0 and 1 since they are tiled. + for (size_t i = 2; i < Rank; i++) { + dst += dst_strides[i] * ptrdiff_t(p[i]); + src += src_strides[i] * ptrdiff_t(p[i]); + } + + // Load a TileSize x TileSize tile into shared memory. Consecutive threads along x + // read consecutive elements along axis 0, keeping the source access coalesced. +#pragma unroll + for (uint32_t x = 0; x < TileSize; x += BlockSizeX) { +#pragma unroll + for (uint32_t y = 0; y < TileSize; y += BlockSizeY) { + uint32_t tx = threadIdx.x + x; + uint32_t ty = threadIdx.y + y; + uint32_t px = p[0] * TileSize + tx; + uint32_t py = p[1] * TileSize + ty; + T value {}; + + if (px < extent0 && py < extent1) { + value = *reinterpret_cast( + src + ptrdiff_t(px) * src_strides[0] + ptrdiff_t(py) * src_strides[1] + ); + } + + shared_tile[ty][tx] = value; + } + } + + __syncthreads(); + + // Write the tile back transposed. Consecutive threads along x now write consecutive + // elements along axis 1, keeping the destination access coalesced. +#pragma unroll + for (uint32_t x = 0; x < TileSize; x += BlockSizeX) { +#pragma unroll + for (uint32_t y = 0; y < TileSize; y += BlockSizeY) { + uint32_t tx = threadIdx.x + x; + uint32_t ty = threadIdx.y + y; + uint32_t px = p[0] * TileSize + ty; + uint32_t py = p[1] * TileSize + tx; + + if (px < extent0 && py < extent1) { + *reinterpret_cast( + dst + ptrdiff_t(px) * dst_strides[0] + ptrdiff_t(py) * dst_strides[1] + ) = shared_tile[tx][ty]; + } + } + } + + __syncthreads(); + + linear_index += gridDim.x; + } +} + +template +__global__ void elementwise_copy_kernel( + std::byte* dst_addr, + const std::byte* src_addr, + IndexMapper mapper, + Vec dst_strides, + Vec src_strides +) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Vec p; + + while (mapper.unravel(linear_index, p)) { + std::byte* dst = dst_addr; + const std::byte* src = src_addr; + + dst += ptrdiff_t(sizeof(T)) * ptrdiff_t(p[0]); + src += ptrdiff_t(sizeof(T)) * ptrdiff_t(p[0]); + +#pragma unroll + for (size_t i = 1; i < Rank; i++) { + dst += dst_strides[i] * ptrdiff_t(p[i]); + src += src_strides[i] * ptrdiff_t(p[i]); + } + + *reinterpret_cast(dst) = *reinterpret_cast(src); + linear_index += blockDim.x * gridDim.x; + } +} + +//=========================================================================== +// Fill kernels (see fill_gpu.cu) +//=========================================================================== + +template +__global__ void elementwise_fill_kernel( + std::byte* dst_addr, + T value, + IndexMapper mapper, + Vec strides +) { + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + threadIdx.x; + Vec p; + + while (mapper.unravel(linear_index, p)) { + std::byte* addr = dst_addr; + + // strides[0] == sizeof(T) + addr += ptrdiff_t(sizeof(T)) * ptrdiff_t(p[0]); + +#pragma unroll + for (size_t i = 1; i < Rank; i++) { + addr += strides[i] * ptrdiff_t(p[i]); + } + + *reinterpret_cast(addr) = value; + linear_index += blockDim.x * gridDim.x; + } +} + +//=========================================================================== +// Reduction kernels (see reduction_gpu.cu) +//=========================================================================== + +// Wavefront/warp width for CUDA and HIP +#if defined(KMM_USE_HIP) && defined(__HIP_DEVICE_COMPILE__) +constexpr uint32_t KMM_REDUCE_WARP_SIZE = __AMDGCN_WAVEFRONT_SIZE__; +#else +constexpr uint32_t KMM_REDUCE_WARP_SIZE = 32; +#endif + +// Generic warp shuffle that works for any trivially-copyable `T` (not just the scalar types that +// `__shfl_xor_sync` natively overloads, such as `KeyValue<...>`), by shuffling it word-by-word. +template +KMM_DEVICE T shfl_xor(T value, int offset) { + static_assert(sizeof(T) % sizeof(uint32_t) == 0, "size of T must be a multiple of 4 bytes"); + constexpr size_t num_words = sizeof(T) / sizeof(uint32_t); + uint32_t words[num_words]; + +#if defined(KMM_USE_HIP) + // HIP does not consider `std::memcpy` callable from device code, so copy word-by-word instead. +#pragma unroll + for (size_t i = 0; i < num_words; i++) { + words[i] = reinterpret_cast(&value)[i]; + } +#else + std::memcpy(words, &value, sizeof(T)); +#endif + +#pragma unroll + for (auto& word : words) { +#if defined(KMM_USE_HIP) + word = __shfl_xor(word, offset); +#else + word = __shfl_xor_sync(0xffffffffu, word, offset); +#endif + } + +#if defined(KMM_USE_HIP) +#pragma unroll + for (size_t i = 0; i < num_words; i++) { + reinterpret_cast(&value)[i] = words[i]; + } +#else + std::memcpy(&value, words, sizeof(T)); +#endif + + return value; +} + +/// Performs `dst_addr[I] += src_addr[I]` for each index I in `mapper`. +template +__global__ void elementwise_fold_kernel( + const std::byte* __restrict__ src_addr, + std::byte* __restrict__ dst_addr, + IndexMapper mapper, + Vec input_strides, + Vec output_strides +) { + using T = typename Reduction::element_type; + uint32_t tx = threadIdx.x; + uint32_t linear_index = uint32_t(blockIdx.x) * BlockSize + tx; + Vec p; + + while (mapper.unravel(linear_index, p)) { + const auto* src = src_addr; + auto* dst = dst_addr; + +#pragma unroll + for (size_t i = 0; i < N; i++) { + src += input_strides[i] * ptrdiff_t(p[i]); + dst += output_strides[i] * ptrdiff_t(p[i]); + } + + Reduction accum = Reduction {*reinterpret_cast(dst)}; + + T value = *reinterpret_cast(src); + accum.consume(value); + + *reinterpret_cast(dst) = accum.finish(); + linear_index += BlockSize * gridDim.x; + } +} + +/// Indicates what kind of reduction must be performed. In all cases: +/// * blockDim.x: number of outputs produced by one block. +/// * blockDim.y: number of threads that cooperate to produce that one output. +enum class ReductionKernelKind { + // kernel is launched with thread block dimensions (1, BlockSize). Each thread is given a unique output and + // will perform the full reduction for that one output. + Elementwise, + + // kernel is launched with thread block dimensions (WARP_SIZE, BlockSize/WARP_SIZE). Each thread in a warp + // is given a different output, but there are multiple warps that will work together to perform the reduction. + Warpwise, + + // kernel is launched with thread block dimensions (BlockSize, 1). All threads are given the same output and they + // must all work together to perform the reduction. + Blockwise +}; + +template +__global__ void elementwise_reduce_kernel( + const std::byte* __restrict__ src_addr, + std::byte* __restrict__ dst_addr, + bool accumulate, + uint32_t reduction_extent, + ptrdiff_t reduction_stride, + IndexMapper mapper, + Vec input_strides, + Vec output_strides, + ReductionKernelKind kind +) { + using T = typename Reduction::element_type; + uint32_t tx = threadIdx.x; + __shared__ T shared_values[2][BlockSize / 2]; + + uint32_t linear_index = uint32_t(blockIdx.x) * blockDim.x + tx; + Vec p {}; + + for (;; linear_index += blockDim.x * gridDim.x) { + // from some reason, using `int` instead of `bool` leads to much cleaner ptx. + int active = mapper.unravel(linear_index, p); + + // only if all threads in this block want to exit, we actually exit. This prevents deadlock on the + // upcoming __syncthreads as all threads must remain active for that to work. + if (!__syncthreads_or(active)) { + break; + } + + const auto* src = src_addr + ptrdiff_t(threadIdx.y) * reduction_stride; + auto* dst = dst_addr; + +#pragma unroll + for (size_t i = 0; i < N; i++) { + src += input_strides[i] * ptrdiff_t(p[i]); + dst += output_strides[i] * ptrdiff_t(p[i]); + } + + Reduction accum {}; + + if (active) { + for (uint32_t current = threadIdx.y; current < reduction_extent; + current += blockDim.y) { + T value = *reinterpret_cast(src); + accum.consume(value); + src += reduction_stride * blockDim.y; + } + } + + // if not elementwise (i.e., one thread per output), then we must do a reduction + if (kind != ReductionKernelKind::Elementwise) { + bool parity = false; + uint32_t tid = threadIdx.y * blockDim.x + threadIdx.x; + + // reduce the values across the block into a single warp. +#pragma unroll + for (uint32_t stride = BlockSize / 2; stride >= KMM_REDUCE_WARP_SIZE; stride /= 2) { + parity = !parity; + + // threads that are active and tid >= stride (or alternatively, having NOT tid < stride), + // will write there value and then deactivate themselves. + if (active && tid >= stride) { + // the __syncthreads_or protects the first write, the successive writes will have + // __syncthreads from each iteration to protect shared memory. + shared_values[parity][tid - stride] = accum.finish(); + active = false; + } + + __syncthreads(); + + if (active) { + // the above __syncthreads ensures that all values have been written. + accum.consume(shared_values[parity][tid]); + } + } + + // if blockwise, we must reduce the warp to a single value. + if (kind == ReductionKernelKind::Blockwise) { +#pragma unroll + for (uint32_t offset = KMM_REDUCE_WARP_SIZE / 2; offset >= 1; offset /= 2) { + accum.consume(shfl_xor(accum.finish(), int(offset))); + } + } + + // Only threads in the first warp remain active + active &= threadIdx.y == 0; + } + + if (active) { + if (accumulate) { + accum.consume(*reinterpret_cast(dst)); + } + + *reinterpret_cast(dst) = accum.finish(); + } + } +} + +} // namespace kmm::memops diff --git a/src/runtime/memops/reduction.cpp b/src/runtime/memops/reduction.cpp new file mode 100644 index 00000000..fb48b3dd --- /dev/null +++ b/src/runtime/memops/reduction.cpp @@ -0,0 +1,416 @@ +#include +#include +#include + +#include "simplify_dims.hpp" + +#include "kmm/core/const_value.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/copy.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/reducer.hpp" +#include "kmm/runtime/memops/reduction.hpp" + +namespace kmm { + +ReductionDescription ReductionDescription::simplify() const { + ReductionDescription result = *this; + + result.num_dims = simplify_dims( + dims, + num_dims, + result.dims, + [](const ReductionDim& a, const ReductionDim& b) { + return unsigned_abs(a.input_stride) > unsigned_abs(b.input_stride); + }, + [](ReductionDim& outer, const ReductionDim& inner) { + if (outer.input_stride == inner.input_stride * inner.extent + && outer.output_stride == inner.output_stride * inner.extent) { + outer.extent *= inner.extent; + outer.input_stride = inner.input_stride; + outer.output_stride = inner.output_stride; + return true; + } + + return false; + } + ); + + return result; +} + +// Widens the half-open byte range [lo, hi) so that visiting `extent` positions spaced `stride` +// bytes apart (starting anywhere in the current range) stays inside it. +static void extend_range( + ptrdiff_t& lo, + ptrdiff_t& hi, + memops_extent_type extent, + memops_stride_type stride +) { + ptrdiff_t span = checked_mul(extent - 1, stride); + lo = span < 0 ? checked_add(lo, span) : lo; + hi = span > 0 ? checked_add(hi, span) : hi; +} + +// The byte range touched relative to `base_offset` when visiting every batch position (`dims`, +// via `stride_member`) and, for each, every reduced position (`reduced_extent`/`reduced_stride`). +// An axis of extent zero means nothing is visited, so the range collapses to empty. +static Range offset_range( + ptrdiff_t base_offset, + const ReductionDim* dims, + size_t num_dims, + size_t element_size, + memops_stride_type ReductionDim::* stride_member, + memops_extent_type reduced_extent, + memops_stride_type reduced_stride +) { + for (size_t i = 0; i < num_dims; i++) { + if (dims[i].extent < 1) { + return {base_offset, base_offset}; + } + } + + if (reduced_extent < 1) { + return {base_offset, base_offset}; + } + + ptrdiff_t lo = base_offset; + ptrdiff_t hi = base_offset; + + for (size_t i = 0; i < num_dims; i++) { + extend_range(lo, hi, dims[i].extent, dims[i].*stride_member); + } + + extend_range(lo, hi, reduced_extent, reduced_stride); + + return {lo, checked_add(hi, element_size)}; +} + +Range ReductionDescription::src_range() const { + return offset_range( + static_cast(input_offset), + dims, + num_dims, + data_type_size(dtype), + &ReductionDim::input_stride, + reduction_extent, + reduction_stride + ); +} + +Range ReductionDescription::dst_range() const { + return offset_range( + static_cast(output_offset), + dims, + num_dims, + data_type_size(dtype), + &ReductionDim::output_stride, + /* reduced_extent = */ 1, + /* reduced_stride = */ 0 + ); +} + +CopyDescription ReductionDescription::as_copy() const { + KMM_ASSERT(is_equivalent_to_copy()); + + CopyDescription result(data_type_size(dtype)); + result.src_offset = input_offset; + result.dst_offset = output_offset; + result.num_dims = num_dims; + + for (size_t i = 0; i < num_dims; i++) { + result.dims[i] = CopyDim {dims[i].extent, dims[i].input_stride, dims[i].output_stride}; + } + + return result; +} + +FillDescription ReductionDescription::as_fill() const { + KMM_ASSERT(is_equivalent_to_fill()); + + FillDescription result(reduction_identity(dtype, operation)); + result.offset = output_offset; + result.num_dims = num_dims; + + for (size_t i = 0; i < num_dims; i++) { + result.dims[i] = FillDim {dims[i].extent, dims[i].output_stride}; + } + + return result; +} + +/// Returns the identity for `T`/`Op` when that is a supported reduction (see +/// `is_reduction_supported`), otherwise throws instead of failing to compile. +template +static FillValue reduction_identity_checked() { + if constexpr (memops::is_reduction_supported) { + return FillValue::from(memops::Reducer {}.finish()); + } else { + throw std::runtime_error("reduction operator not supported for this data type"); + } +} + +template +static FillValue reduction_identity_typed(ReductionOp op) { + switch (op) { + case ReductionOp::Sum: + return reduction_identity_checked(); + case ReductionOp::Product: + return reduction_identity_checked(); + case ReductionOp::Min: + return reduction_identity_checked(); + case ReductionOp::Max: + return reduction_identity_checked(); + case ReductionOp::BitwiseAnd: + return reduction_identity_checked(); + case ReductionOp::BitwiseOr: + return reduction_identity_checked(); + } + + KMM_PANIC("invalid reduction operator"); +} + +FillValue reduction_identity(DataType dtype, ReductionOp op) { + switch (dtype) { + case DataType::Int32: + return reduction_identity_typed(op); + case DataType::Int64: + return reduction_identity_typed(op); + case DataType::Uint32: + return reduction_identity_typed(op); + case DataType::Uint64: + return reduction_identity_typed(op); + case DataType::Float32: + return reduction_identity_typed(op); + case DataType::Float64: + return reduction_identity_typed(op); + case DataType::KeyValueInt64: + return reduction_identity_typed>(op); + case DataType::KeyValueFloat64: + return reduction_identity_typed>(op); + case DataType::Unknown: + break; + } + + KMM_PANIC("invalid data type"); +} + +template +static void reduce_leaf( + const std::byte* src, + std::byte* dst, + E reduction_extent, + memops_stride_type reduction_stride +) { + // `src`/`dst` and every stride are assumed to be `T`-aligned. This holds for any description + // built by `make_reduction_description`, which derives every offset and stride from + // `element_size == sizeof(T)`. + using T = typename Reduction::element_type; + Reduction acc {}; + + if constexpr (Accumulate) { + T previous = *reinterpret_cast(dst); + acc = Reduction {previous}; + } + + for (memops_extent_type i = 0; i < reduction_extent; i++) { + T value = *reinterpret_cast(src + i * reduction_stride); + acc.consume(value); + } + + *reinterpret_cast(dst) = acc.finish(); +} + +/// Recurses over the batch axes with `Rank` (the number of remaining axes) as a template +/// parameter, so the compiler can fully unroll the loop nest for the common, small ranks instead +/// of looping over a runtime-sized `dims` array. +template +static void reduce_dim( + const std::byte* src, + std::byte* dst, + const ReductionDim* dims, + E reduction_extent, + memops_stride_type reduction_stride +) { + if constexpr (Rank == 0) { + reduce_leaf(src, dst, reduction_extent, reduction_stride); + } else { + for (memops_extent_type i = 0; i < dims->extent; i++) { + reduce_dim( + src + i * dims->input_stride, + dst + i * dims->output_stride, + dims + 1, + reduction_extent, + reduction_stride + ); + } + } +} + +/// Dispatches the runtime `num_dims` (at most `MEMOPS_MAX_DIMS`, checked by the caller) to the +/// matching `reduce_dim` instantiation. +template +static void reduce_dim_dispatch( + size_t num_dims, + const std::byte* src, + std::byte* dst, + const ReductionDim* dims, + E reduction_extent, + memops_stride_type reduction_stride +) { + if (num_dims == Rank) { + reduce_dim(src, dst, dims, reduction_extent, reduction_stride); + } else if constexpr (Rank > 0) { + reduce_dim_dispatch( + num_dims, + src, + dst, + dims, + reduction_extent, + reduction_stride + ); + } else { + KMM_PANIC("invalid number of dimensions"); + } +} + +template +static void reduce_op_accumulate( + const void* src_addr, + void* dst_addr, + const ReductionDescription& description +) { + // cases are: + // * reduction_extent == 1: simply reduce src into dst + // * reduction_extent == 2: reduce two buffers into one + // * reduction_extent > 2: arbitrary reduction axis + if (description.reduction_extent == 1) { + reduce_dim_dispatch( + description.num_dims, + static_cast(src_addr) + description.input_offset, + static_cast(dst_addr) + description.output_offset, + description.dims, + ConstValue(), + description.reduction_stride + ); + } else if (description.reduction_extent == 2) { + reduce_dim_dispatch( + description.num_dims, + static_cast(src_addr) + description.input_offset, + static_cast(dst_addr) + description.output_offset, + description.dims, + ConstValue(), + description.reduction_stride + ); + } else { + reduce_dim_dispatch( + description.num_dims, + static_cast(src_addr) + description.input_offset, + static_cast(dst_addr) + description.output_offset, + description.dims, + description.reduction_extent, + description.reduction_stride + ); + } +} + +template +static void reduce_op( + const void* src_addr, + void* dst_addr, + const ReductionDescription& description +) { + if (description.accumulate) { + reduce_op_accumulate(src_addr, dst_addr, description); + } else { + reduce_op_accumulate(src_addr, dst_addr, description); + } +} + +/// Runs `reduce_op` when that combination is a supported reduction (see +/// `is_reduction_supported`), otherwise throws instead of failing to compile. This is what lets +/// `reduce_typed` list every operator unconditionally even for element types that only support a +/// subset (integers-only for the bitwise ops, `KeyValue` only for `Min`/`Max`). +template +static void reduce_op_checked( + const void* src_addr, + void* dst_addr, + const ReductionDescription& description +) { + if constexpr (memops::is_reduction_supported) { + reduce_op>(src_addr, dst_addr, description); + } else { + throw std::runtime_error("reduction operator not supported for this data type"); + } +} + +template +static void reduce_typed( + const void* src_addr, + void* dst_addr, + const ReductionDescription& description +) { + switch (description.operation) { + case ReductionOp::Sum: + return reduce_op_checked(src_addr, dst_addr, description); + case ReductionOp::Product: + return reduce_op_checked(src_addr, dst_addr, description); + case ReductionOp::Min: + return reduce_op_checked(src_addr, dst_addr, description); + case ReductionOp::Max: + return reduce_op_checked(src_addr, dst_addr, description); + case ReductionOp::BitwiseAnd: + return reduce_op_checked(src_addr, dst_addr, description); + case ReductionOp::BitwiseOr: + return reduce_op_checked(src_addr, dst_addr, description); + } + + KMM_PANIC("invalid reduction operator"); +} + +namespace memops { + +void reduce(const void* src_addr, void* dst_addr, const ReductionDescription& description) { + if (description.is_noop()) { + return; + } + + if (description.is_equivalent_to_copy()) { + return copy(src_addr, dst_addr, description.as_copy()); + } + + if (description.is_equivalent_to_fill()) { + return fill(dst_addr, description.as_fill()); + } + + switch (description.dtype) { + case DataType::Unknown: + break; + case DataType::Int32: + return reduce_typed(src_addr, dst_addr, description); + case DataType::Int64: + return reduce_typed(src_addr, dst_addr, description); + case DataType::Uint32: + return reduce_typed(src_addr, dst_addr, description); + case DataType::Uint64: + return reduce_typed(src_addr, dst_addr, description); + case DataType::Float32: + return reduce_typed(src_addr, dst_addr, description); + case DataType::Float64: + return reduce_typed(src_addr, dst_addr, description); + case DataType::KeyValueInt64: + return reduce_typed>(src_addr, dst_addr, description); + case DataType::KeyValueFloat64: + return reduce_typed>(src_addr, dst_addr, description); + } + + KMM_PANIC("invalid data type"); +} + +} // namespace memops + +// `memops::reduce_gpu` (the GPU counterpart) lives in reduction_gpu.cu -- it needs real GPU kernels +// (via CUB), so it must be compiled by nvcc, unlike the rest of this file. + +} // namespace kmm diff --git a/src/runtime/memops/reduction_gpu.cu b/src/runtime/memops/reduction_gpu.cu new file mode 100644 index 00000000..8bcdf033 --- /dev/null +++ b/src/runtime/memops/reduction_gpu.cu @@ -0,0 +1,590 @@ +#include +#include +#include +#include +#include + +#include "memops_gpu_kernels.cuh" + +#include "kmm/core/checked_compare.hpp" +#include "kmm/core/fast_divisor.hpp" +#include "kmm/core/integer_fun.hpp" +#include "kmm/core/vec.hpp" +#include "kmm/runtime/memops/copy_gpu.hpp" +#include "kmm/runtime/memops/fill_gpu.hpp" +#include "kmm/runtime/memops/reducer.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/memops/reduction_gpu.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm::memops { + +struct ReductionGPU: ReductionDescription { + g_stream_t stream; + const void* src_addr; + void* dst_addr; + + template + bool is_aligned() const { + if (!is_divisible(reinterpret_cast(src_addr), alignof(T))) { + return false; + } + + if (!is_divisible(reinterpret_cast(dst_addr), alignof(T))) { + return false; + } + + if (!is_divisible(reduction_stride, alignof(T))) { + return false; + } + + for (size_t i = 0; i < num_dims; i++) { + if (!is_divisible(dims[i].input_stride, alignof(T))) { + return false; + } + + if (!is_divisible(dims[i].output_stride, alignof(T))) { + return false; + } + } + + return true; + } +}; + +ReductionGPU make_plan( + g_stream_t stream, + const void* src_base, + void* dst_base, + const ReductionDescription& description +) { + ReductionGPU plan; + plan.dtype = description.dtype; + plan.operation = description.operation; + plan.stream = stream; + plan.src_addr = static_cast(src_base) + description.input_offset; + plan.dst_addr = static_cast(dst_base) + description.output_offset; + plan.accumulate = description.accumulate; + plan.reduction_stride = description.reduction_stride; + plan.reduction_extent = description.reduction_extent > 0 ? description.reduction_extent : 0; + plan.num_dims = 0; + + size_t old_rank = description.num_dims; + size_t new_rank = 0; + ReductionDim dims[MEMOPS_MAX_DIMS] = {}; + std::copy_n(description.dims, old_rank, dims); + + // Sorted by ascending `input_stride` (rather than `output_stride`): the axis at index 0 is + // read repeatedly (once per reduction step, in every kernel that consults it for coalescing), + // while its output element is written only once, so minimizing the input-side stride at index + // 0 matters more than minimizing the output-side one. + for (size_t i = 0; i < old_rank; i++) { + for (size_t j = i + 1; j < old_rank; j++) { + if (unsigned_abs(dims[j].input_stride) < unsigned_abs(dims[i].input_stride)) { + std::swap(dims[i], dims[j]); + } + } + + auto new_dim = dims[i]; + + // no reduction needed, just say there is one dim with extent==0 + if (new_dim.extent <= 0) { + plan.dims[0].extent = 0; + plan.dims[0].input_stride = 0; + plan.dims[0].output_stride = 0; + new_rank = 1; + break; + } + + // skip this dimension if it has extent of one + if (new_dim.extent == 1) { + continue; + } + + // if the dst_stride is zero, all values land at the same location. We can effectively consider its + // extent to be equal to one. + if (new_dim.output_stride == 0) { + // TODO: Maybe throw an exception here? Why would you want a dst stride of zero? + continue; + } + + // fix negative stride by subtracting offset from the pointer + if (new_dim.output_stride < 0) { + plan.src_addr = static_cast(plan.src_addr) + + (new_dim.extent - 1) * new_dim.input_stride; + plan.dst_addr = static_cast(plan.dst_addr) + + (new_dim.extent - 1) * new_dim.output_stride; + + new_dim.output_stride = -new_dim.output_stride; + new_dim.input_stride = -new_dim.input_stride; + } + + if (new_rank > 0) { + auto k = new_rank - 1; + + if (is_equal(new_dim.output_stride, plan.dims[k].output_stride * plan.dims[k].extent) + && is_equal( + new_dim.input_stride, + plan.dims[k].input_stride * plan.dims[k].extent + )) { + plan.dims[new_rank - 1].extent *= new_dim.extent; + continue; + } + + if (is_divisible(new_dim.output_stride, plan.dims[k].output_stride) + && is_divisible(new_dim.input_stride, plan.dims[k].input_stride)) { + // TODO: should be something smart when the strides are multiples of each other + } + } + + plan.dims[new_rank] = new_dim; + new_rank++; + } + + plan.num_dims = new_rank; + return plan; +} + +// Picks which of the three reduction kernels to launch for one chunked reduction. +ReductionKernelKind select_reduction_kernel_kind( + memops_extent_type reduction_extent, + ptrdiff_t reduction_stride, + ptrdiff_t input_stride0, + ptrdiff_t output_stride0, + size_t element_size +) { + static constexpr int cooperative_threshold = 256; + static constexpr size_t nearly_contiguous_factor = 4; + + if (reduction_extent < cooperative_threshold) { + return ReductionKernelKind::Elementwise; + } + + size_t nearly_contiguous_bound = nearly_contiguous_factor * element_size; + bool reduction_axis_contiguous = unsigned_abs(reduction_stride) <= nearly_contiguous_bound; + bool output_axis_contiguous = unsigned_abs(input_stride0) <= nearly_contiguous_bound + && unsigned_abs(output_stride0) <= nearly_contiguous_bound; + + if (reduction_axis_contiguous && !output_axis_contiguous) { + return ReductionKernelKind::Blockwise; + } + + return ReductionKernelKind::Warpwise; +} + +template +void launch_reduction_kernel_rank_recur( + g_stream_t stream, + const std::byte* src_addr, + std::byte* dst_addr, + bool accumulate, + memops_extent_type reduction_extent, + ptrdiff_t reduction_stride, + Vec extents, + Vec input_strides, + Vec output_strides +) { + uint32_t num_outputs = 1; + Vec region; + + for (size_t i = 0; i < N; i++) { + if (extents[i] <= 0) { + return; + } + + uint32_t chunk = IndexMapper::max_volume / num_outputs; + + while (extents[i] > chunk) { + Vec head = extents; + head[i] = chunk; + + launch_reduction_kernel_rank_recur( + stream, + src_addr, + dst_addr, + accumulate, + reduction_extent, + reduction_stride, + head, + input_strides, + output_strides + ); + + src_addr += chunk * input_strides[i]; + dst_addr += chunk * output_strides[i]; + extents[i] -= chunk; + } + + region[i] = static_cast(extents[i]); // safe since 0 < extent[i] <= max_volume + num_outputs *= region[i]; + } + + auto mapper = IndexMapper(region); + + // A common case is that one buffer is 'folded' into another buffer. This means that for the parameters + // N=1, accumulate=true, reduction_extent=1 we have a special case kernel. + if constexpr (N == 1) { + if (accumulate && reduction_extent == 1) { + static constexpr uint32_t block_size = 256; + static constexpr uint32_t max_blocks = 4096; + uint32_t grid_size = std::min(max_blocks, div_ceil(num_outputs, block_size)); + + elementwise_fold_kernel<<>>( + src_addr, + dst_addr, + mapper, + input_strides, + output_strides + ); + return; + } + } + + size_t element_size = sizeof(typename Reduction::element_type); + ReductionKernelKind kernel_kind = select_reduction_kernel_kind( + reduction_extent, + reduction_stride, + N > 0 ? input_strides[0] : ptrdiff_t(element_size), + N > 0 ? output_strides[0] : ptrdiff_t(element_size), + element_size + ); + + static constexpr uint32_t threads_per_block = 256; + static constexpr uint32_t max_blocks = 4096; + uint32_t outputs_per_block; + + if (kernel_kind == ReductionKernelKind::Elementwise) { + outputs_per_block = threads_per_block; + } else if (kernel_kind == ReductionKernelKind::Warpwise) { + outputs_per_block = KMM_REDUCE_WARP_SIZE; + } else { + outputs_per_block = 1; + } + + uint32_t reduce_rows = threads_per_block / outputs_per_block; + + // `elementwise_reduce_kernel`'s cross-row combine folds pairs of rows by powers of two, so + // `outputs_per_block` (== `blockDim.x`) must be a power of two dividing `threads_per_block`. + KMM_ASSERT(is_power_of_two(outputs_per_block)); + KMM_ASSERT(outputs_per_block * reduce_rows == threads_per_block); + + dim3 block_size = {outputs_per_block, reduce_rows}; + dim3 grid_size = std::min(max_blocks, div_ceil(num_outputs, outputs_per_block)); + + elementwise_reduce_kernel + <<>>( + src_addr, + dst_addr, + accumulate, + checked_cast(reduction_extent), + reduction_stride, + mapper, + input_strides, + output_strides, + kernel_kind + ); +} + +template +void launch_reduction_kernel_rank(const ReductionGPU& plan) { + KMM_ASSERT(plan.num_dims == N); + Vec extents; + Vec input_strides; + Vec output_strides; + + for (size_t i = 0; i < N; i++) { + extents[i] = checked_cast(plan.dims[i].extent); + input_strides[i] = plan.dims[i].input_stride; + output_strides[i] = plan.dims[i].output_stride; + } + + launch_reduction_kernel_rank_recur( + plan.stream, + reinterpret_cast(plan.src_addr), + reinterpret_cast(plan.dst_addr), + plan.accumulate, + plan.reduction_extent, + plan.reduction_stride, + extents, + input_strides, + output_strides + ); +} + +template +void launch_reduction_kernel_op(const ReductionGPU& plan) { + using T = element_type_t; + size_t num_dims = plan.num_dims; + + if constexpr (!is_reduction_supported) { + throw std::runtime_error( + "invalid reduction parameters: unsupported operation for data type" + ); + } else { + if (!plan.is_aligned()) { + throw std::runtime_error( + "invalid reduction parameters: address not aligned for data type" + ); + } + + if (num_dims == 0) { + ReductionGPU p = plan; + p.num_dims = 1; + p.dims[0] = ReductionDim {}; + launch_reduction_kernel_rank, 1>(p); + } else if (num_dims == 1) { + launch_reduction_kernel_rank, 1>(plan); + } else if (num_dims == 2) { + launch_reduction_kernel_rank, 2>(plan); + } else if (num_dims == 3) { + launch_reduction_kernel_rank, 3>(plan); + } else { + throw std::runtime_error("dimensionality of reduction is too high"); + } + } +} + +template +void launch_reduction_kernel_typed(const ReductionGPU& plan, ReductionOp op) { + static constexpr DataType unsigned_dtype = // + dtype == DataType::Int32 ? DataType::Uint32 + : (dtype == DataType::Int64 ? DataType::Uint64 : dtype); + + switch (op) { + case ReductionOp::Product: + return launch_reduction_kernel_op(plan); + case ReductionOp::Min: + return launch_reduction_kernel_op(plan); + case ReductionOp::Max: + return launch_reduction_kernel_op(plan); + // for these operations, the operation is equivalent on unsigned and signed integers. We use the unsigned + // dtype if possible to minimize the number of reduction kernels that are generated. + case ReductionOp::Sum: + return launch_reduction_kernel_op(plan); + case ReductionOp::BitwiseAnd: + return launch_reduction_kernel_op(plan); + case ReductionOp::BitwiseOr: + return launch_reduction_kernel_op(plan); + default: + throw std::runtime_error("invalid operation for reduction"); + } +} + +void launch_reduction_kernel(const ReductionGPU& plan, DataType dtype, ReductionOp op) { + switch (dtype) { + case DataType::Int32: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::Int64: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::Uint32: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::Uint64: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::Float32: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::Float64: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::KeyValueInt64: + launch_reduction_kernel_typed(plan, op); + break; + case DataType::KeyValueFloat64: + launch_reduction_kernel_typed(plan, op); + break; + default: + throw std::runtime_error("invalid data type for reduction"); + } +} + +constexpr uint32_t min_blocks_for_full_occupancy = 2048; +constexpr size_t reduction_scratch_budget = 16 * 1024 * 1024; +constexpr int32_t min_items_per_chunk = 256; + +void launch_multilevel_reduction( + const ReductionGPU& plan, + int64_t num_chunks, + size_t num_outputs, + void* scratch_addr +) { + KMM_ASSERT(scratch_addr != nullptr); + + size_t element_size = data_type_size(plan.dtype); + size_t num_dims = plan.num_dims; + + // Round the chunk size *down* so that `items_per_chunk * num_chunks <= reduction_extent` + auto items_per_chunk = plan.reduction_extent / num_chunks; + auto tail_extent = plan.reduction_extent - items_per_chunk * num_chunks; + + // Dense (contiguous) strides used to lay out `num_chunks` partial results per output element in scratch. + ptrdiff_t dense_strides[MEMOPS_MAX_DIMS] = {}; + ptrdiff_t dense_stride = ptrdiff_t(element_size); + + for (size_t i = 0; i < num_dims; i++) { + dense_strides[i] = dense_stride; + dense_stride *= plan.dims[i].extent; + } + + ptrdiff_t chunk_output_stride = checked_mul(num_outputs, element_size); + ptrdiff_t chunk_input_stride = checked_mul(items_per_chunk, plan.reduction_stride); + + // Kernel 1: reduce all `num_chunks` equally-sized chunks (each `items_per_chunk` elements) into + // `scratch_addr`, adding the chunk index as an extra batch axis. + { + size_t chunk_pos = num_dims; + + // find the location to inject the chunk_input_stride + for (size_t i = 0; i < num_dims; i++) { + if (unsigned_abs(chunk_input_stride) < unsigned_abs(plan.dims[i].input_stride)) { + chunk_pos = i; + break; + } + } + + ReductionGPU chunk_plan; + chunk_plan.stream = plan.stream; + chunk_plan.src_addr = plan.src_addr; + chunk_plan.dst_addr = scratch_addr; + chunk_plan.accumulate = false; + chunk_plan.reduction_extent = items_per_chunk; + chunk_plan.reduction_stride = plan.reduction_stride; + chunk_plan.num_dims = num_dims + 1; + + for (size_t i = 0; i < num_dims; i++) { + size_t pos = i < chunk_pos ? i : i + 1; + chunk_plan.dims[pos].extent = plan.dims[i].extent; + chunk_plan.dims[pos].input_stride = plan.dims[i].input_stride; + chunk_plan.dims[pos].output_stride = dense_strides[i]; + } + + chunk_plan.dims[chunk_pos].extent = num_chunks; + chunk_plan.dims[chunk_pos].input_stride = chunk_input_stride; + chunk_plan.dims[chunk_pos].output_stride = chunk_output_stride; + + launch_reduction_kernel(chunk_plan, plan.dtype, plan.operation); + } + + // Kernel 2: reduce the leftover tail (`tail_extent < num_chunks` elements). + if (tail_extent > 0) { + ReductionGPU remainder_plan; + remainder_plan.stream = plan.stream; + remainder_plan.src_addr = static_cast(plan.src_addr) + + checked_mul(items_per_chunk, num_chunks) * plan.reduction_stride; + remainder_plan.dst_addr = plan.dst_addr; + remainder_plan.accumulate = plan.accumulate; + remainder_plan.reduction_extent = tail_extent; + remainder_plan.reduction_stride = plan.reduction_stride; + remainder_plan.num_dims = num_dims; + + for (size_t i = 0; i < num_dims; i++) { + remainder_plan.dims[i] = plan.dims[i]; + } + + launch_reduction_kernel(remainder_plan, plan.dtype, plan.operation); + } + + // Kernel 3: fold the `num_chunks` partial results in `scratch_addr` into the real output. When + // Kernel 2 ran it already seeded `dst` (with the caller's `accumulate`), so accumulate on top + // of it here; otherwise this pass is what honors the caller's `accumulate`. + { + ReductionGPU final_plan; + final_plan.stream = plan.stream; + final_plan.src_addr = scratch_addr; + final_plan.dst_addr = plan.dst_addr; + final_plan.accumulate = tail_extent > 0 ? true : plan.accumulate; + final_plan.reduction_extent = num_chunks; + final_plan.reduction_stride = chunk_output_stride; + final_plan.num_dims = num_dims; + + for (size_t i = 0; i < num_dims; i++) { + final_plan.dims[i].extent = plan.dims[i].extent; + final_plan.dims[i].input_stride = dense_strides[i]; + final_plan.dims[i].output_stride = plan.dims[i].output_stride; + } + + launch_reduction_kernel(final_plan, plan.dtype, plan.operation); + } +} + +// Conservative occupancy estimate: assumes the worst case of one block per output (as +// `blockwise_reduce_kernel` does), since that's the regime where extra chunks matter most - +// `elementwise_reduce_kernel`/`warpwise_reduce_kernel` already get plenty of blocks from the +// output dimension alone whenever `num_outputs` is large. +memops_extent_type plan_reduction_chunks( + memops_extent_type reduction_extent, + size_t num_outputs, + size_t element_size +) { + auto max_chunks_by_extent = reduction_extent / min_items_per_chunk; + + if (num_outputs >= min_blocks_for_full_occupancy || max_chunks_by_extent <= 1 + || num_outputs == 0) { + return 1; + } + + memops_extent_type wanted = + div_ceil(min_blocks_for_full_occupancy, num_outputs); + memops_extent_type budget_limit = + reduction_scratch_budget / std::max(num_outputs * element_size, 1); + + return std::max(1, std::min({wanted, budget_limit, max_chunks_by_extent})); +} + +void reduce_gpu( + g_stream_t stream, + const void* src_base, + void* dst_base, + void* scratch_addr, + const ReductionDescription& description +) { + auto simplified = description; + + if (simplified.is_noop()) { + return; + } + + if (simplified.is_equivalent_to_copy()) { + return copy_gpu(stream, src_base, dst_base, simplified.as_copy()); + } + + if (simplified.is_equivalent_to_fill()) { + return fill_gpu(stream, dst_base, simplified.as_fill()); + } + + auto plan = make_plan(stream, src_base, dst_base, simplified); + + size_t element_size = data_type_size(simplified.dtype); + size_t num_outputs = 1; + + for (size_t i = 0; i < plan.num_dims; i++) { + num_outputs = checked_mul(num_outputs, plan.dims[i].extent); + } + + auto num_chunks = plan_reduction_chunks(plan.reduction_extent, num_outputs, element_size); + + if (num_chunks <= 1 || plan.num_dims == MEMOPS_MAX_DIMS) { + launch_reduction_kernel(plan, plan.dtype, plan.operation); + } else { + launch_multilevel_reduction(plan, num_chunks, num_outputs, scratch_addr); + } +} + +size_t reduce_gpu_scratch_size(const ReductionDescription& description) { + size_t element_size = data_type_size(description.dtype); + auto num_outputs = description.num_outputs(); + auto reduction_extent = std::max(description.reduction_extent, 0); + + auto num_chunks = plan_reduction_chunks(reduction_extent, num_outputs, element_size); + + if (num_chunks <= 1) { + return 0; + } + + return checked_mul(num_chunks, num_outputs) * element_size; +} + +} // namespace kmm::memops diff --git a/src/runtime/memops/simplify_dims.hpp b/src/runtime/memops/simplify_dims.hpp new file mode 100644 index 00000000..b2673e91 --- /dev/null +++ b/src/runtime/memops/simplify_dims.hpp @@ -0,0 +1,84 @@ +#pragma once + +#include +#include + +#include "kmm/runtime/memops/types.hpp" + +namespace kmm { + +/// Shared core of `CopyDescription::simplify`, `FillDescription::simplify`, and +/// `ReductionDescription::simplify`. Normalizes the axis list into one of exactly two canonical +/// forms: +/// +/// - *empty* (some axis has extent zero or negative): a single axis with `extent == 0` and every +/// stride zero. The offending axis's real strides are discarded, so every empty description +/// simplifies to the same value regardless of which axis vanished. +/// - *non-empty*: axes with extent one are dropped (they are visited exactly once, so their +/// stride never contributes to addressing), the rest are sorted into descending stride order +/// using `less`, and adjacent axes are merged whenever `try_merge` reports that the outer axis +/// simply repeats the inner axis's memory layout without gaps. The result has zero axes (a +/// single element) or only axes with `extent >= 2`, no two of which are mergeable. +/// +/// So `extent < 0` and `extent == 1` never survive `simplify`, and `extent == 0` survives only as +/// the single all-zero sentinel axis. +/// +/// The extent-zero case is handled up front, rather than folded into `try_merge` like the +/// contiguous-merge case, for two reasons: there may be no preceding axis yet to fold it into +/// (an all-zero-then-nonzero input has no "outer" for the first axis), and `try_merge` only ever +/// looks at stride-adjacent neighbors, so it can't guarantee collapsing to a single axis when the +/// zero-extent axis isn't adjacent (in sorted order) to every other axis. Callers rely on that +/// single-axis guarantee to hit their fast paths for an empty operation instead of falling back +/// to their "unsupported layout" case. +/// +/// `less` must be a strict weak ordering over `Dim` that sorts by descending "primary" stride +/// (with whatever tie-breaking a given description wants). `try_merge(outer, inner)` must fold +/// `inner` into `outer` and return `true` if they are contiguous, or return `false` (leaving +/// `outer` untouched) otherwise. `out` must have room for `num_dims` entries. +template +size_t simplify_dims(const Dim* dims, size_t num_dims, Dim* out, Less less, TryMerge try_merge) { + size_t n = 0; + + for (size_t i = 0; i < num_dims; i++) { + Dim dim = dims[i]; + + // A zero (or negative, which describes the same nothing) extent makes the whole + // operation empty. Collapse to one canonical all-zero axis. + if (dim.extent <= 0) { + out[0] = Dim {}; + out[0].extent = 0; + return 1; + } + + // Drop axes of extent one: visited exactly once, so their stride never affects addressing. + if (dim.extent == 1) { + continue; + } + + out[n] = dim; + n++; + } + + // bubble sort + for (size_t i = 0; i < MEMOPS_MAX_DIMS; i++) { + for (size_t j = 0; j < i; j++) { + if (i < n && less(out[i], out[j])) { + std::swap(out[i], out[j]); + } + } + } + + // Compact in place: `num_out` never exceeds `i`, so writing `out[num_out]` never overwrites + size_t num_out = 0; + + for (size_t i = 0; i < n; i++) { + if (num_out == 0 || !try_merge(out[num_out - 1], out[i])) { + out[num_out] = out[i]; + num_out++; + } + } + + return num_out; +} + +} // namespace kmm diff --git a/src/runtime/memops/types.cpp b/src/runtime/memops/types.cpp new file mode 100644 index 00000000..a5c4c089 --- /dev/null +++ b/src/runtime/memops/types.cpp @@ -0,0 +1,71 @@ +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memops/types.hpp" + +namespace kmm { + +size_t data_type_size(DataType dtype) { + switch (dtype) { + case DataType::Unknown: + break; + case DataType::Int32: + case DataType::Uint32: + case DataType::Float32: + return 4; + case DataType::Int64: + case DataType::Uint64: + case DataType::Float64: + return 8; + case DataType::KeyValueInt64: + return sizeof(KeyValue); + case DataType::KeyValueFloat64: + return sizeof(KeyValue); + } + + KMM_PANIC("invalid data type"); +} + +const char* data_type_name(DataType dtype) { + switch (dtype) { + case DataType::Unknown: + break; + case DataType::Int32: + return "Int32"; + case DataType::Int64: + return "Int64"; + case DataType::Uint32: + return "Uint32"; + case DataType::Uint64: + return "Uint64"; + case DataType::Float32: + return "Float32"; + case DataType::Float64: + return "Float64"; + case DataType::KeyValueInt64: + return "KeyValueInt64"; + case DataType::KeyValueFloat64: + return "KeyValueFloat64"; + } + + KMM_PANIC("invalid data type"); +} + +const char* reduction_op_name(ReductionOp op) { + switch (op) { + case ReductionOp::Sum: + return "Sum"; + case ReductionOp::Product: + return "Product"; + case ReductionOp::Min: + return "Min"; + case ReductionOp::Max: + return "Max"; + case ReductionOp::BitwiseAnd: + return "BitwiseAnd"; + case ReductionOp::BitwiseOr: + return "BitwiseOr"; + } + + KMM_PANIC("invalid reduction operator"); +} + +} // namespace kmm diff --git a/src/runtime/memory_buffer.cpp b/src/runtime/memory_buffer.cpp new file mode 100644 index 00000000..cffb4f9f --- /dev/null +++ b/src/runtime/memory_buffer.cpp @@ -0,0 +1,499 @@ +#include "fmt/chrono.h" +#include "spdlog/spdlog.h" + +#include "kmm/runtime/memory_buffer.hpp" + +namespace kmm { + +Poll HostAccessControl::poll_pending_future() { + if (pending_future.valid()) { + if (pending_future.wait_for(std::chrono::seconds(0)) == std::future_status::timeout) { + return Poll::Pending; + } + + // will not block but does clear the future and handle any exceptions + wait_pending_future(); + } + + return Poll::Ready; +} + +void HostAccessControl::wait_pending_future() { + if (pending_future.valid()) { + auto before = std::chrono::system_clock::now(); + + try { + pending_future.get(); + } catch (...) { + pending_future = {}; + is_valid = false; + throw; + } + + auto after = std::chrono::system_clock::now(); + auto duration = after - before; + + if (duration > std::chrono::milliseconds(1)) { + spdlog::warn("waited for {} for host before to become available", duration); + } + } +} + +bool MemoryBufferImpl::is_compatible(MemoryId memory_id, AccessKind mode) noexcept { + for (auto r = queue_head; r != nullptr; r = r->queue_next) { + // two exclusive access are never allowed + if (r->mode == AccessKind::Exclusive || mode == AccessKind::Exclusive) { + return false; + } + + // if one writes, then they must access the same memory + if (r->mode != AccessKind::ReadOnly || mode != AccessKind::ReadOnly) { + if (r->memory_id != memory_id) { + return false; + } + } + } + + // if we reach this point, then access has been granted + return true; +} + +bool MemoryBufferImpl::try_register_request(BufferQueueNode* req) noexcept { + if (released) { + return false; + } + + if (!is_compatible(req->memory_id, req->mode)) { + return false; + } + + req->queue_prev = queue_tail; + req->queue_next = nullptr; + + if (queue_tail != nullptr) { + queue_tail->queue_next = req; + } else { + queue_head = req; + } + + queue_tail = req; + return true; +} + +void MemoryBufferImpl::unregister_request(BufferQueueNode* req) noexcept { + if (req->queue_prev != nullptr) { + req->queue_prev->queue_next = req->queue_next; + } else { + queue_head = req->queue_next; + } + + if (req->queue_next != nullptr) { + req->queue_next->queue_prev = req->queue_prev; + } else { + queue_tail = req->queue_prev; + } + + req->queue_prev = nullptr; + req->queue_next = nullptr; +} + +AllocResult MemoryBufferImpl::try_allocate_location( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId dst_id +) { + auto& dst_loc = location(dst_id); + KMM_ASSERT(!dst_loc.is_allocated); + + MemoryId src_id = dst_id; + bool has_peer = find_valid_location(dst_id, src_id); + + // A host operation migt still be active. Calls treat returning `false` as out of memory + // failure. Instead, we fall back to a regular allocation and let the caller ensure that + // the data is copied when the future completes. + bool peer_ready = !src_id.is_host() || host_location.poll_pending_future() == Poll::Ready; + + if (has_peer && peer_ready + && (dst_id.is_host() || src_id.is_host() + || data->is_copy_supported(system, src_id, dst_id))) { + auto& src_loc = location(src_id); + DeviceEventSet events; + AllocResult result = data->allocate_and_copy( + system, + src_id, + dst_id, + stream_hint, + src_loc.retrieve_access(AccessKind::SharedWrite), + events + ); + + if (result == AllocResult::Success) { + src_loc.record_access(AccessKind::ReadOnly, events); + dst_loc.mark_allocated_and_valid(events); + } + + return result; + } + + DeviceEventSet deps; + AllocResult result = data->allocate(system, dst_id, stream_hint, deps); + + if (result == AllocResult::Success) { + dst_loc.mark_allocated(std::move(deps)); + } + + return result; +} + +bool MemoryBufferImpl::allocate_host(MemorySystem& system, const DeviceStreamId& stream_hint) { + spdlog::debug("allocate buffer {} in memory {}", name, MemoryId::host()); + + if (try_allocate_location(system, stream_hint, MemoryId::host()) != AllocResult::Success) { + throw std::runtime_error("could not allocate, out of host memory"); + } + + return true; +} + +void MemoryBufferImpl::increment_host_users() noexcept { + spdlog::trace( + "buffer {}: host alloc_count {} -> {}", + name, + host_location.alloc_count, + host_location.alloc_count + 1 + ); + host_location.alloc_count++; +} + +void MemoryBufferImpl::decrement_host_users() noexcept { + spdlog::trace( + "buffer {}: host alloc_count {} -> {}", + name, + host_location.alloc_count, + host_location.alloc_count - 1 + ); + KMM_ASSERT(host_location.alloc_count > 0); + host_location.alloc_count--; +} + +bool MemoryBufferImpl::deallocate_host(MemorySystem& system, const DeviceStreamId& stream_hint) { + auto& loc = host_location; + KMM_ASSERT(loc.alloc_count == 0); + + if (!is_allocated(MemoryId::host())) { + return false; + } + + // This is the one place that actually frees the host allocation, so unlike + // `invalidate_other_allocs`/`invalidate_all` (which may leave a stale future draining in + // the background) it must force the location to finish. + loc.wait_pending_future(); + + spdlog::debug("deallocate buffer {} in memory {}", name, MemoryId::host()); + + auto deps = loc.mark_deallocated(); + data->deallocate(system, MemoryId::host(), stream_hint, std::move(deps)); + return true; +} + +AllocResult MemoryBufferImpl::try_allocate_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId id +) { + spdlog::debug("allocate buffer {} in memory {}", name, MemoryId::device(id)); + return try_allocate_location(system, stream_hint, MemoryId::device(id)); +} + +bool MemoryBufferImpl::deallocate_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId id, + DeviceLRU& lru +) { + auto& loc = device_locations[id.get()]; + KMM_ASSERT(loc.alloc_count == 0); + + if (!is_allocated(MemoryId::device(id))) { + return false; + } + + spdlog::debug("deallocate buffer {} in memory {}", name, MemoryId::device(id)); + + auto deps = loc.mark_deallocated(); + data->deallocate(system, MemoryId::device(id), stream_hint, std::move(deps)); + lru.remove(&loc); + + return true; +} + +void MemoryBufferImpl::increment_device_users(DeviceId id, DeviceLRU& lru) noexcept { + auto& loc = device_locations[id.get()]; + spdlog::trace( + "buffer {}: device {} alloc_count {} -> {}", + name, + id, + loc.alloc_count, + loc.alloc_count + 1 + ); + loc.alloc_count++; + + if (loc.alloc_count == 1) { + lru.remove(&loc); + } +} + +void MemoryBufferImpl::decrement_device_users(DeviceId id, DeviceLRU& lru) noexcept { + auto& loc = device_locations[id.get()]; + KMM_ASSERT(loc.alloc_count > 0); + spdlog::trace( + "buffer {}: device {} alloc_count {} -> {}", + name, + id, + loc.alloc_count, + loc.alloc_count - 1 + ); + loc.alloc_count--; + + if (loc.alloc_count == 0 && evictable) { + lru.insert(&loc); + } +} + +void MemoryBufferImpl::evict_device( + MemorySystem& system, + const DeviceStreamId& stream_hint, + DeviceId memory_id, + DeviceLRU& lru +) { + auto& loc = device_locations[memory_id.get()]; + KMM_ASSERT(is_allocated(MemoryId::device(memory_id))); + + auto other_id = MemoryId::host(); + + // if this entry is valid and there are no other valid entries, then we must evict to host + if (loc.is_valid && !find_valid_location(MemoryId::device(memory_id), other_id)) { + // `allocate_host` finds this device (still valid, not yet deallocated) as its copy + // source and copies from it directly -- src is a device, so this can never be Pending. + allocate_host(system, stream_hint); + } + + deallocate_device(system, stream_hint, memory_id, lru); +} + +Poll MemoryBufferImpl::ensure_alloc_valid( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId memory_id +) { + auto& loc = location(memory_id); + + // If we are accessing the host, we must first check if the future on the host is ready. + // This is needed for both the case that is is_valid==true and is_valid==false. + if (memory_id.is_host() && host_location.poll_pending_future() == Poll::Pending) { + return Poll::Pending; + } + + // if already valid, we just exit. Even for the host, there can not be a future pending. + if (loc.is_valid) { + return Poll::Ready; + } + + MemoryId peer_id = memory_id; + bool has_valid_peer = find_valid_location(memory_id, peer_id); + + if (!has_valid_peer) { + // + const auto& deps = location(memory_id).retrieve_access(AccessKind::Exclusive); + + // This now becomes the home memory + if (!home_memory_id.has_value()) { + home_memory_id = memory_id; + } + + spdlog::debug("initializing buffer {} on {}", name, memory_id); + + if (memory_id.is_host()) { + host_location.pending_future = data->initialize_host(system, deps); + loc.mark_valid(DeviceEvent::null()); + return host_location.poll_pending_future(); + } else { + auto event = data->initialize_device(system, memory_id.as_device(), stream_hint, deps); + loc.mark_valid(event); + return Poll::Ready; + } + } else if (memory_id.is_host() || peer_id.is_host() + || data->is_copy_supported(system, peer_id, memory_id)) { + // copy D2H or H2D or D2D (if possible) + return poll_copy(system, stream_hint, peer_id, memory_id); + } else { + // copy D2H -> H2D: `allocate_host` finds `peer_id` (or another valid location) as its + // copy source and performs the D2H leg itself; only the H2D leg remains here. + allocate_host(system, stream_hint); + do_copy(system, stream_hint, MemoryId::host(), memory_id); + return Poll::Ready; + } +} + +void MemoryBufferImpl::invalidate_other_allocs(MemoryId memory_id) { + DeviceEventSet deps; + + if (memory_id.is_host()) { + // Invalidate all device entries + for (size_t i = 0; i < MAX_DEVICES; i++) { + auto& peer_entry = device_locations[i]; + deps.insert(peer_entry.retrieve_access(AccessKind::ReadOnly)); + peer_entry.is_valid = false; + } + } else { + // Invalidate host if necessary. We deliberately do NOT wait for a still-running host + // future here: the underlying allocation stays put (only `deallocate_host` frees it, + // and it force-drains any in-flight future first, so there's no use-after-free risk), + // and the only other hazard -- a later host fill overwriting `pending_future` while + // this one is still outstanding -- is guarded non-blockingly in `ensure_alloc_valid`. + deps.insert(host_location.retrieve_access(AccessKind::ReadOnly)); + host_location.is_valid = false; + + // Invalidate all _other_ device entries + for (size_t i = 0; i < MAX_DEVICES; i++) { + if (memory_id == MemoryId::device(DeviceId(i))) { + continue; + } + + auto& peer_entry = device_locations[i]; + deps.insert(peer_entry.retrieve_access(AccessKind::ReadOnly)); + peer_entry.is_valid = false; + } + } + + location(memory_id).record_access(AccessKind::Exclusive, deps); +} + +DeviceEventSet MemoryBufferImpl::invalidate_all() { + if (queue_head != nullptr) { + throw std::runtime_error("failed to invalidate buffer as access is locked by a request"); + } + + DeviceEventSet deps; + + // See `invalidate_other_allocs`: no need to wait out an in-flight host future here -- + // the allocation isn't freed by this call, and `ensure_alloc_valid` non-blockingly drains + // any leftover future before a new one could overwrite it. + deps.insert(host_location.retrieve_access(AccessKind::ReadOnly)); + host_location.is_valid = false; + + for (size_t i = 0; i < MAX_DEVICES; i++) { + auto& peer_entry = device_locations[i]; + deps.insert(peer_entry.retrieve_access(AccessKind::ReadOnly)); + peer_entry.is_valid = false; + } + + return deps; +} + +// wait for: +// Exclusive waits for Exclusive, SharedWrite, ReadOnly +// ReadOnly wait for Exclusive, SharedWrite +// SharedWrite wait for Exclusive +static AccessKind waiting_mode_for(AccessKind mode) { + if (mode == AccessKind::Exclusive) { + return AccessKind::ReadOnly; + } else if (mode == AccessKind::ReadOnly) { + return AccessKind::SharedWrite; + } else { + return AccessKind::Exclusive; + } +} + +Poll MemoryBufferImpl::before_access( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId memory_id, + AccessKind mode +) { + // 1) ensure that the allocation contains valid data + if (ensure_alloc_valid(system, stream_hint, memory_id) == Poll::Pending) { + return Poll::Pending; + } + + // 2) if we are going to write, we must invalidate all the others + if (mode != AccessKind::ReadOnly) { + invalidate_other_allocs(memory_id); + } + + // Hint the eventual access so the backend can prefetch, if it wants to. + auto& loc = location(memory_id); + data->hint_access(system, memory_id, stream_hint, loc.retrieve_access(waiting_mode_for(mode))); + + return Poll::Ready; +} + +BufferAccessor MemoryBufferImpl::access( + MemoryId memory_id, + AccessKind mode, + DeviceEventSet& deps_out +) { + auto& loc = location(memory_id); + deps_out.insert(loc.retrieve_access(waiting_mode_for(mode))); + + return BufferAccessor { + memory_id, + size_in_bytes, + mode != AccessKind::ReadOnly, + data->address(memory_id) + }; +} + +void MemoryBufferImpl::after_access( + MemoryId memory_id, + AccessKind mode, + const DeviceEventSet& deps +) { + location(memory_id).record_access(mode, deps); +} + +Poll MemoryBufferImpl::poll_copy( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId src_id, + MemoryId dst_id +) { + // if this involves the host, we must first wait until the associated future completes. If + // the future is still active, then there may still be threads actively reading/writing + // to the host memory and we must wait until they complete. + if ((src_id.is_host() || dst_id.is_host()) + && host_location.poll_pending_future() == Poll::Pending) { + return Poll::Pending; + } + + do_copy(system, stream_hint, src_id, dst_id); + return Poll::Ready; +} + +void MemoryBufferImpl::do_copy( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryId src_id, + MemoryId dst_id +) { + auto& src_alloc = location(src_id); + auto& dst_alloc = location(dst_id); + + KMM_ASSERT(src_alloc.is_allocated && dst_alloc.is_allocated); + KMM_ASSERT(src_alloc.is_valid && !dst_alloc.is_valid); + + // if this involves the host, then there cannot be any futures active on the host memory. + KMM_ASSERT((!src_id.is_host() && !dst_id.is_host()) || !host_location.pending_future.valid()); + + DeviceEventSet deps; + deps.insert(src_alloc.retrieve_access(AccessKind::SharedWrite)); + deps.insert(dst_alloc.retrieve_access(AccessKind::ReadOnly)); + + spdlog::debug("launch copy for buffer {} from {} to {} (deps: {})", name, src_id, dst_id, deps); + data->copy(system, src_id, dst_id, stream_hint, deps, deps); + + src_alloc.record_access(AccessKind::ReadOnly, deps); + dst_alloc.mark_valid(deps); +} + +} // namespace kmm \ No newline at end of file diff --git a/src/runtime/memory_manager.cpp b/src/runtime/memory_manager.cpp index 49e77303..0101e725 100644 --- a/src/runtime/memory_manager.cpp +++ b/src/runtime/memory_manager.cpp @@ -1,1142 +1,770 @@ -#include -#include +#include +#include +#include +#include +#include -#include "fmt/format.h" +#include "fmt/chrono.h" #include "spdlog/spdlog.h" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/memory_buffer.hpp" #include "kmm/runtime/memory_manager.hpp" -#include "kmm/utils/integer_fun.hpp" namespace kmm { -struct BufferEntry { - KMM_NOT_COPYABLE_OR_MOVABLE(BufferEntry) - - public: - BufferEntry() = default; - - bool is_allocated = false; - bool is_valid = false; - - size_t num_allocation_locks = 0; - - // let `<=` the happens-before operator. Then it should ALWAYS hold that - // * epoch_event <= write_events - // * write_events <= access_events - DeviceEventSet epoch_event; - DeviceEventSet write_events; - DeviceEventSet access_events; -}; - -struct HostEntry: public BufferEntry { - void* data = nullptr; -}; - -struct DeviceEntry: public BufferEntry { - g_device_ptr_t data = 0; +struct MemoryTransactionImpl: reference_count { + explicit MemoryTransactionImpl(uint64_t id, MemoryTransaction parent) : + id(id), + parent(std::move(parent)) {} - MemoryManager::Buffer* lru_older = nullptr; - MemoryManager::Buffer* lru_newer = nullptr; + const uint64_t id; + MemoryTransaction parent; // may be null }; -struct MemoryManager::Transaction { - KMM_NOT_COPYABLE_OR_MOVABLE(Transaction); +struct DeviceQueueNode { + DeviceQueueNode(NotifyHandle callback) : callback(std::move(callback)) {} - public: - Transaction(uint64_t id, std::shared_ptr parent) : id(id), parent(parent) {} + NotifyHandle callback; - uint64_t id; - std::shared_ptr parent; - std::chrono::system_clock::time_point created_at = std::chrono::system_clock::now(); + // Intrusive doubly-linked list used by `DeviceState`: a single list holding + // both the allocated prefix and the waiting suffix, split by `req_first_pending`. + DeviceQueueNode* device_prev = nullptr; + DeviceQueueNode* device_next = nullptr; + MemoryTransaction parent; }; -struct MemoryManager::Request { - KMM_NOT_COPYABLE_OR_MOVABLE(Request); +struct MemoryRequestImpl: BufferQueueNode, DeviceQueueNode, reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(MemoryRequestImpl) public: - Request( - uint64_t identifier, - std::shared_ptr buffer, + MemoryRequestImpl( + uint64_t id, + refcnt_ptr buffer, MemoryId memory_id, - AccessMode mode, - std::shared_ptr parent + AccessKind mode, + NotifyHandle callback ) : - identifier(identifier), - buffer(std::move(buffer)), - memory_id(memory_id), - mode(mode), - parent(std::move(parent)) {} + BufferQueueNode(memory_id, mode, callback), + DeviceQueueNode(callback), + id(id), + buffer(std::move(buffer)) {} + + enum struct State { Unqueued, WaitingForAllocation, Granted, Ready, Released }; - enum struct Status { Init, Allocated, Locked, Ready, Deleted }; - Status status = Status::Init; - uint64_t identifier; - std::shared_ptr buffer; - MemoryId memory_id; - AccessMode mode; - std::shared_ptr parent; - std::chrono::system_clock::time_point created_at = std::chrono::system_clock::now(); - - bool allocation_acquired = false; - Request* allocation_next = nullptr; - Request* allocation_prev = nullptr; - - bool access_acquired = false; - Request* access_next = nullptr; - Request* access_prev = nullptr; + const uint64_t id; + const refcnt_ptr buffer; + State state = State::Unqueued; }; -struct MemoryManager::Buffer { - KMM_NOT_COPYABLE_OR_MOVABLE(Buffer); +// Recovers the `MemoryBufferImpl` owning `device_locations[id.get()]` from a +// pointer to that location, using the fact that it lives at a fixed offset +// inside its owner (the LRU list only ever holds device locations, never +// `host_location`, so `id` is enough to identify which slot `loc` is). +static MemoryBufferImpl* owner_of_device_location(DeviceAccessControl* loc, DeviceId id) noexcept { + DeviceAccessControl* first = loc - id.get(); + + // `MemoryBufferImpl` derives from `reference_count`, so it is not standard-layout and + // `offsetof` is technically only conditionally-supported here. The offset is still valid in + // practice (single, non-virtual inheritance), so silence the warning rather than the check. +#if defined(__GNUC__) || defined(__clang__) + #pragma GCC diagnostic push + #pragma GCC diagnostic ignored "-Winvalid-offsetof" +#endif + auto offset = offsetof(MemoryBufferImpl, device_locations); +#if defined(__GNUC__) || defined(__clang__) + #pragma GCC diagnostic pop +#endif + + return reinterpret_cast(reinterpret_cast(first) - offset); +} - public: - Buffer(std::string name, BufferLayout layout) : name(std::move(name)), layout(layout) { - if (this->name.empty()) { - this->name = std::to_string(std::intptr_t(this)); +// Host-memory counterpart of `DeviceState`: there is only ever one host +// location per buffer, allocation never blocks or needs to evict to make +// room, so there is no LRU and no request queue here, just byte accounting. +struct HostState { + // Request-scoped acquire: bundles allocation (if needed) with the usage-count + // bump. Its counterpart `release_for_request` only undoes the usage-count + // half — deallocation is buffer-scoped, not request-scoped, see `deallocate`. + void acquire_for_request( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryBufferImpl* buf + ) { + if (!buf->is_allocated(MemoryId::host())) { + buf->allocate_host(system, stream_hint); + bytes_allocated += buf->size_in_bytes; } - } - std::string name; - BufferLayout layout; - HostEntry host_entry; - DeviceEntry device_entry[MAX_DEVICES]; - size_t num_requests_active = 0; - bool is_deleted = false; + buf->increment_host_users(); + } - Request* access_head = nullptr; - Request* access_first_pending = nullptr; - Request* access_tail = nullptr; + void release_for_request(MemoryBufferImpl* buf) noexcept { + buf->decrement_host_users(); + } - KMM_INLINE BufferEntry& entry(MemoryId memory_id) noexcept { - if (memory_id.is_host()) { - return host_entry; - } else { - return device_entry[memory_id.as_device()]; + void deallocate( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryBufferImpl* buf + ) noexcept { + if (buf->deallocate_host(system, stream_hint)) { + bytes_allocated -= buf->size_in_bytes; } } - void add_to_access_queue(Request& req) noexcept { - this->num_requests_active++; + // Total number of bytes currently allocated on the host across all buffers. + size_t bytes_allocated = 0; +}; - if (this->access_tail == nullptr) { - this->access_head = &req; - this->access_tail = &req; - } else { - auto* prev = this->access_tail; - prev->access_next = &req; - req.access_prev = prev; +struct DeviceState { + DeviceState(DeviceId id) noexcept : memory_id(id) {} - this->access_tail = &req; - } + void register_request(DeviceQueueNode* req) noexcept { + req->device_prev = req_tail; + req->device_next = nullptr; - if (this->access_first_pending == nullptr) { - this->access_first_pending = &req; + if (req_tail != nullptr) { + req_tail->device_next = req; + } else { + req_head = req; } - } - void remove_from_access_queue(Request& req) noexcept { - this->num_requests_active--; - auto* prev = std::exchange(req.access_prev, nullptr); - auto* next = std::exchange(req.access_next, nullptr); + req_tail = req; - if (prev != nullptr) { - prev->access_next = next; + if (req_first_pending == nullptr) { + req_first_pending = req; } else { - KMM_ASSERT(this->access_head == &req); - this->access_head = next; + // This request joins behind an already-waiting `req_first_pending`. + // Its transaction's ancestor chain may reach a transaction that + // `is_out_of_memory` previously judged not blocked (e.g. if this + // request's transaction is a descendant of one already holding + // memory here), which would flip that earlier conclusion. Wake + // `req_first_pending` so it rechecks instead of waiting forever + // on a now-stale answer. + req_first_pending->callback.notify(); } + } - if (next != nullptr) { - next->access_prev = prev; + void unregister_request(DeviceQueueNode* req) noexcept { + if (req->device_prev != nullptr) { + req->device_prev->device_next = req->device_next; } else { - KMM_ASSERT(this->access_tail == &req); - this->access_tail = prev; + req_head = req->device_next; } - if (this->access_first_pending == &req) { - this->access_first_pending = next; + if (req->device_next != nullptr) { + req->device_next->device_prev = req->device_prev; + } else { + req_tail = req->device_prev; } - if (req.access_acquired) { - req.access_acquired = false; - - // Poll queue, releasing the lock might allow another request to gain access - poll_access_queue(); - } - } + if (req_first_pending == req) { + req_first_pending = req->device_next; - bool is_access_allowed(const Request& req) const noexcept { - auto mode = req.mode; - - for (auto* it = this->access_head; it != this->access_first_pending; it = it->access_next) { - // Two exclusive requests can never be granted access simultaneously - if (mode == AccessMode::Exclusive || it->mode == AccessMode::Exclusive) { - return false; - } - - // Two non-read requests can only be granted simultaneously if operating on the same memory. - if (mode != AccessMode::Read || it->mode != AccessMode::Read) { - if (req.memory_id != it->memory_id) { - return false; - } + // another request became the head of the pending queue, notify it + if (req_first_pending != nullptr) { + req_first_pending->callback.notify(); } } - return true; + req->callback.clear(); + req->device_prev = nullptr; + req->device_next = nullptr; } - void poll_access_queue() noexcept { - while (this->access_first_pending != nullptr) { - auto* req = this->access_first_pending; + // Counterpart of `try_acquire_for_request` for release: only undoes the + // usage-count half of acquire, since deallocation is buffer-scoped, not + // request-scoped (see `deallocate`, only called from buffer release). + void release_for_request(MemoryBufferImpl* buf) noexcept { + buf->decrement_device_users(memory_id, lru); - if (!is_access_allowed(*req)) { - return; - } - - req->access_acquired = true; - spdlog::trace( - "access to buffer {} was granted to request {} (memory={}, mode={})", - this->name, - req->identifier, - req->memory_id, - req->mode - ); - - this->access_first_pending = this->access_first_pending->access_next; + // if somebody is waiting, wake them up! We must always wake up the pending request, + // even if we do not add anything to the LRU, since the request needs to check again + // if the may be out of memory now. + if (req_first_pending != nullptr) { + req_first_pending->callback.notify(); } } - bool request_has_access(const Request& req) { - poll_access_queue(); - return req.access_acquired; - } -}; - -struct MemoryManager::Device { - KMM_NOT_COPYABLE_OR_MOVABLE(Device); - - public: - DeviceId device_id; - bool printed_offload_warning = false; - - Buffer* lru_oldest = nullptr; - Buffer* lru_newest = nullptr; + AllocResult try_allocate( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryBufferImpl* buf + ) { + if (!buf->is_allocated(MemoryId::device(memory_id))) { + auto result = buf->try_allocate_device(system, stream_hint, memory_id); - Request* allocation_head = nullptr; - Request* allocation_first_pending = nullptr; - Request* allocation_tail = nullptr; - - Device(DeviceId id) : device_id(id) {} - - void add_to_allocation_queue(Request& req) noexcept { - auto* tail = this->allocation_tail; + // could not allocate, exit now. + if (result != AllocResult::Success) { + return result; + } - if (tail == nullptr) { - this->allocation_head = &req; - } else { - tail->allocation_next = &req; - req.allocation_prev = tail; + // increment byte count. + bytes_allocated += buf->size_in_bytes; } - this->allocation_tail = &req; - - if (this->allocation_first_pending == nullptr) { - this->allocation_first_pending = &req; - } + // increment user count. + buf->increment_device_users(memory_id, lru); + return AllocResult::Success; } - void remove_from_allocation_queue(Request& req) noexcept { - auto* prev = std::exchange(req.allocation_prev, nullptr); - auto* next = std::exchange(req.allocation_next, nullptr); - - if (prev != nullptr) { - prev->allocation_next = next; - } else { - KMM_ASSERT(this->allocation_head == &req); - this->allocation_head = next; - } - - if (next != nullptr) { - next->allocation_prev = prev; - } else { - KMM_ASSERT(this->allocation_tail == &req); - this->allocation_tail = prev; - } - - if (this->allocation_first_pending == &req) { - this->allocation_first_pending = next; + void deallocate( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryBufferImpl* buf + ) noexcept { + if (buf->is_allocated(MemoryId::device(memory_id))) { + buf->deallocate_device(system, stream_hint, memory_id, lru); + bytes_allocated -= buf->size_in_bytes; } } - void add_to_lru(Buffer& buffer) noexcept { - auto& device_entry = buffer.device_entry[device_id]; - auto& device = *this; + MemoryBufferImpl* select_evict_victim() { + // First, find a buffer which has a different home location than memory_id + for (auto* node = lru.least_recently_used(); node != nullptr; node = node->lru_next) { + // This is done in a very hacky way by using some pointer trickery. + MemoryBufferImpl* owner = owner_of_device_location(node, memory_id); - KMM_ASSERT(device_entry.is_allocated); - KMM_ASSERT(device_entry.num_allocation_locks == 0); + if (owner->home_memory_id != MemoryId::device(memory_id)) { + return owner; + } + } - auto* prev = device.lru_newest; - if (prev != nullptr) { - prev->device_entry[device_id].lru_newer = &buffer; - } else { - device.lru_oldest = &buffer; + // Second, just select the most recently used + if (auto* victim = lru.least_recently_used()) { + return owner_of_device_location(victim, memory_id); } - device_entry.lru_older = prev; - device_entry.lru_newer = nullptr; - device.lru_newest = &buffer; + return nullptr; } - void remove_from_lru(Buffer& buffer) noexcept { - auto& device_entry = buffer.device_entry[device_id]; - auto& device = *this; - - KMM_ASSERT(device_entry.is_allocated); - KMM_ASSERT(device_entry.num_allocation_locks == 0); - - auto* prev = std::exchange(device_entry.lru_newer, nullptr); - auto* next = std::exchange(device_entry.lru_older, nullptr); - - if (prev != nullptr) { - prev->device_entry[device_id].lru_older = next; + bool try_evict_one(MemorySystem& system, const DeviceStreamId& stream_hint) { + if (auto* victim = select_evict_victim()) { + spdlog::debug("out of memory on {}, evicting buffer {}", memory_id, victim->name); + return try_evict_buffer(system, stream_hint, victim); } else { - device.lru_newest = next; - } - - if (next != nullptr) { - next->device_entry[device_id].lru_newer = prev; - } else { - device.lru_oldest = prev; + spdlog::debug("out of memory on {}, no buffer available for eviction", memory_id); + return false; } } - void increment_allocation_locks(Buffer& buffer, Request& req) { - KMM_ASSERT(req.allocation_acquired == false); - KMM_ASSERT(allocation_first_pending == &req); - - auto& device_entry = buffer.device_entry[device_id]; - KMM_ASSERT(device_entry.is_allocated); + bool try_evict_buffer( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryBufferImpl* buf + ) { + auto& loc = buf->device_locations[memory_id.get()]; - if (device_entry.num_allocation_locks == 0) { - remove_from_lru(buffer); + if (!loc.in_lru) { + return false; } - device_entry.num_allocation_locks++; - - req.allocation_acquired = true; - this->allocation_first_pending = req.allocation_next; + buf->evict_device(system, stream_hint, memory_id, lru); + bytes_allocated -= buf->size_in_bytes; + spdlog::debug("buffer {} has been evicted from {}", buf->name, memory_id); + return true; } - void decrement_allocation_locks(Buffer& buffer) { - auto& device_entry = buffer.device_entry[device_id]; - KMM_ASSERT(device_entry.is_allocated); - KMM_ASSERT(device_entry.num_allocation_locks > 0); - - device_entry.num_allocation_locks--; - - if (device_entry.num_allocation_locks == 0) { - add_to_lru(buffer); - } - } + // Defined below, once `MemoryRequestImpl` is a complete type. + bool try_acquire_for_request( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryRequestImpl* req + ); + bool is_out_of_memory(MemoryRequestImpl* req); + + DeviceId memory_id; + DeviceLRU lru; + bool has_printed_oom_warning = false; + + // Total number of bytes currently allocated on this device across all buffers + // (i.e. the sum of `nbytes` of every buffer with an allocated `Location` here). + size_t bytes_allocated = 0; + + /// This is a linked list of requests wanting to allocate memory. + /// - req_head to req_first_pending: all requests that have been assigned memory + /// - req_first_pending to req_tail: all requests that are still waiting for memory + DeviceQueueNode* req_head = nullptr; + DeviceQueueNode* req_first_pending = nullptr; + DeviceQueueNode* req_tail = nullptr; }; -template -static auto make_devices(std::index_sequence /*unused*/) { - return new MemoryManager::Device[MAX_DEVICES] {DeviceId(Id)...}; -} - -MemoryManager::MemoryManager(std::shared_ptr memory_system) : - m_memory(std::move(memory_system)), - m_devices(make_devices(std::make_index_sequence())) {} - -MemoryManager::~MemoryManager() { - KMM_ASSERT(m_buffers.empty()); -} - -bool MemoryManager::is_idle(DeviceStreamManager& streams) const { - bool result = true; - - for (const auto& buffer : m_buffers) { - for (auto& e : buffer->device_entry) { - result &= streams.is_ready(e.access_events); - result &= streams.is_ready(e.write_events); - result &= streams.is_ready(e.epoch_event); - } - - auto& e = buffer->host_entry; - result &= streams.is_ready(e.access_events); - result &= streams.is_ready(e.write_events); - result &= streams.is_ready(e.epoch_event); - } - - return result; -} - -std::shared_ptr MemoryManager::create_transaction( - std::shared_ptr parent -) { - auto id = m_next_transaction_id++; - return std::make_shared(id, parent); -} - -std::shared_ptr MemoryManager::create_buffer( - BufferLayout layout, - std::string name +bool DeviceState::try_acquire_for_request( + MemorySystem& system, + const DeviceStreamId& stream_hint, + MemoryRequestImpl* req ) { - // Size cannot be zero - if (layout.size_in_bytes == 0) { - layout.size_in_bytes = 1; - } - - // Make sure alignment is power of two and size is multiple of alignment - layout.alignment = round_up_to_power_of_two(layout.alignment); - layout.size_in_bytes = round_up_to_multiple(layout.size_in_bytes, layout.alignment); + KMM_ASSERT(req_first_pending != nullptr); - auto buffer = std::make_shared(std::move(name), std::move(layout)); - m_buffers.emplace(buffer); - return buffer; -} - -void MemoryManager::delete_buffer(std::shared_ptr buffer) { - if (buffer->is_deleted) { - return; - } - - KMM_ASSERT(buffer->num_requests_active == 0); - KMM_ASSERT(buffer->access_head == nullptr); - KMM_ASSERT(buffer->access_first_pending == nullptr); - KMM_ASSERT(buffer->access_tail == nullptr); - - buffer->is_deleted = true; - m_buffers.erase(buffer); - - deallocate_host(*buffer); - - for (size_t i = 0; i < MAX_DEVICES; i++) { - deallocate_device_async(DeviceId(i), *buffer); + // only the first pending request may allocate + if (req_first_pending != req) { + return false; } - check_consistency(); -} + auto* buf = req->buffer.get(); -std::shared_ptr MemoryManager::create_request( - std::shared_ptr buffer, - MemoryId memory_id, - AccessMode mode, - std::shared_ptr parent -) { - KMM_ASSERT(!buffer->is_deleted); - auto req = std::make_shared(m_next_request_id++, buffer, memory_id, mode, parent); + // if it is already allocated, then we are done + while (true) { + auto result = try_allocate(system, stream_hint, buf); - m_active_requests.insert(req); - buffer->add_to_access_queue(*req); + // Success! Return true. + if (result == AllocResult::Success) { + if (has_printed_oom_warning) { + has_printed_oom_warning = false; + spdlog::info("{} is no longer out of memory", memory_id); + } - if (memory_id.is_device()) { - device_at(memory_id.as_device()).add_to_allocation_queue(*req); - } + req->callback.clear(); + req_first_pending = req->device_next; - return req; -} + // notify the next request that it can now attempt allocation + if (req_first_pending != nullptr) { + req_first_pending->callback.notify(); + } -Poll MemoryManager::poll_request(Request& req, DeviceEventSet& deps_out) { - auto& buffer = *req.buffer; - auto memory_id = req.memory_id; + return true; + } - if (req.status == Request::Status::Init) { - if (memory_id.is_host()) { - lock_allocation_host(buffer, memory_id.device_affinity(), req); - } else { - if (!try_lock_allocation_device(memory_id.as_device(), buffer, req)) { - return Poll::Pending; - } + // Allocation failed, we need to retry later + if (result == AllocResult::ErrorPending) { + return false; } - req.status = Request::Status::Allocated; - } + // + if (result != AllocResult::ErrorOutOfMemory) { + auto message = fmt::format( + "allocation failed for device {}: allocation is not supported", + memory_id + ); + spdlog::error("{}", message); + throw std::runtime_error(message); + } - if (req.status == Request::Status::Allocated) { - if (!buffer.request_has_access(req)) { - return Poll::Pending; + if (!has_printed_oom_warning) { + has_printed_oom_warning = true; + spdlog::warn( + "{} is out of memory. The system will now offload unused data from " + "GPU memory to host memory, which may significantly degrade performance.", + memory_id + ); } - req.status = Request::Status::Locked; - } + // Out of memory: reclaim the least-recently-used eligible location for + // this device and retry. If successful, retry + if (try_evict_one(system, stream_hint)) { + continue; + } - if (req.status == Request::Status::Locked) { - prepare_access_to_buffer(memory_id, buffer, req.mode, deps_out); - req.status = Request::Status::Ready; - } + if (is_out_of_memory(req)) { + auto message = fmt::format("allocation failed for device {}: out of memory", memory_id); + spdlog::error("{}", message); + throw std::runtime_error(message); + } - if (req.status == Request::Status::Ready) { - return Poll::Ready; + // Everything failed. return false + return false; } - - throw std::runtime_error("cannot poll a deleted request"); } -void MemoryManager::release_request(std::shared_ptr req, DeviceEvent event) { - auto memory_id = req->memory_id; - auto& buffer = *req->buffer; - auto status = std::exchange(req->status, Request::Status::Deleted); - - if (status == Request::Status::Ready) { - status = Request::Status::Locked; - } - - if (status == Request::Status::Locked) { - spdlog::trace( - "access to buffer {} was revoked from request {} (memory={}, mode={}, GPU event={})", - buffer.name, - req->identifier, - req->memory_id, - req->mode, - event - ); +bool DeviceState::is_out_of_memory(MemoryRequestImpl* req) { + KMM_ASSERT(req_first_pending == req); - finalize_access_to_buffer(memory_id, buffer, req->mode, event); - status = Request::Status::Allocated; - } + ankerl::unordered_dense::set not_blocked; - if (status == Request::Status::Allocated) { - if (memory_id.is_host()) { - unlock_allocation_host(buffer, *req); - } else { - unlock_allocation_device(memory_id.as_device(), buffer, *req); + // Transactions currently holding memory are inserted into not_blocked + for (auto* it = req_head; it != req_first_pending; it = it->device_next) { + for (auto* t = it->parent.get(); t != nullptr; t = t->parent.get()) { + not_blocked.insert(t); } - - status = Request::Status::Init; } - if (status == Request::Status::Init) { - if (memory_id.is_device()) { - device_at(memory_id.as_device()).remove_from_allocation_queue(*req); + // Transactions waiting for more memory are assume to be blocked and are thus removed + // from not_blocked. This gives a list of transactions that have memory assigned and are + // not currently waiting for more memory. + for (auto* it = req_first_pending; it != nullptr; it = it->device_next) { + for (auto* t = it->parent.get(); t != nullptr; t = t->parent.get()) { + not_blocked.erase(t); } - - buffer.remove_from_access_queue(*req); - m_active_requests.erase(req); } -} - -BufferAccessor MemoryManager::get_accessor(Request& req) { - KMM_ASSERT(req.status == Request::Status::Ready); - KMM_ASSERT(m_buffers.count(req.buffer) > 0); - - const auto& buffer = *req.buffer; - void* address; - if (req.memory_id.is_host()) { - address = buffer.host_entry.data; - } else { - address = reinterpret_cast(buffer.device_entry[req.memory_id.as_device()].data); + // Whatever remains is a transaction that holds memory here and isn't + // itself blocked waiting for more, i.e., it can complete on its own. + if (!not_blocked.empty()) { + return false; } - return BufferAccessor { - .memory_id = req.memory_id, - .layout = buffer.layout, - .is_writable = req.mode != AccessMode::Read, - .address = address - }; + // TODO: log warning + return true; } -void MemoryManager::allocate_host(Buffer& buffer, DeviceId device_affinity) { - auto& host_entry = buffer.host_entry; - size_t size_in_bytes = buffer.layout.size_in_bytes; +struct MemoryManager::Impl { + KMM_NOT_COPYABLE_OR_MOVABLE(Impl) - KMM_ASSERT(host_entry.is_allocated == false); - KMM_ASSERT(host_entry.num_allocation_locks == 0); - - spdlog::trace("allocate {} bytes on host", size_in_bytes, buffer.name); + public: + explicit Impl(refcnt_ptr memory_system) : + Impl(std::move(memory_system), std::make_index_sequence()) {} - void* ptr; - DeviceEventSet events; - AllocationResult result = m_memory->allocate_host(size_in_bytes, device_affinity, &ptr, events); + private: + template + Impl(refcnt_ptr memory_system, std::index_sequence) : + memory_system(std::move(memory_system)), + device_states {DeviceState(DeviceId(Is))...} {} - if (result != AllocationResult::Success) { - throw std::runtime_error("could not allocate, out of host memory"); + public: + MemorySystem& system() { + return *memory_system; } - host_entry.data = ptr; - host_entry.is_allocated = true; - host_entry.is_valid = false; - host_entry.epoch_event = events; - host_entry.access_events = events; - host_entry.write_events = std::move(events); -} - -void MemoryManager::deallocate_host(Buffer& buffer) { - auto& host_entry = buffer.host_entry; - if (!host_entry.is_allocated) { - return; + DeviceState& device(DeviceId id) { + return device_states[id.get()]; } - KMM_ASSERT(host_entry.num_allocation_locks == 0); - KMM_ASSERT(buffer.access_head == nullptr); - KMM_ASSERT(buffer.access_first_pending == nullptr); - KMM_ASSERT(buffer.access_tail == nullptr); - - size_t size_in_bytes = buffer.layout.size_in_bytes; - spdlog::trace( - "free {} bytes on host (dependencies={})", - size_in_bytes, - buffer.name, - host_entry.access_events - ); - - m_memory->deallocate_host(host_entry.data, size_in_bytes, std::move(host_entry.access_events)); - - host_entry.data = nullptr; - host_entry.is_allocated = false; - host_entry.is_valid = false; - host_entry.epoch_event.clear(); - host_entry.write_events.clear(); - host_entry.access_events.clear(); - host_entry.data = nullptr; -} - -bool MemoryManager::try_free_device_memory(DeviceId device_id) { - auto& device = device_at(device_id); - auto* victim = device.lru_oldest; - - if (victim == nullptr) { - return false; + HostState& host() { + return host_state; } - if (victim->device_entry[device_id].is_valid) { - bool valid_anywhere = victim->host_entry.is_valid; - + // Frees any device location that's sitting idle in its LRU while marked + // invalid: it holds no useful data, so there is no reason to wait for dealloc. + void reclaim_invalidated(const DeviceStreamId& stream_hint, MemoryBufferImpl* buf) { for (size_t i = 0; i < MAX_DEVICES; i++) { - if (i != device_id) { - valid_anywhere |= victim->device_entry[i].is_valid; - } - } + auto& loc = buf->device_locations[i]; - if (!valid_anywhere) { - if (!device.printed_offload_warning) { - device.printed_offload_warning = true; - spdlog::warn( - "GPU {} is out of memory. The system will now offload data from GPU memory to " - "host memory as a fallback, which may significantly reduce performance.", - device_id - ); + if (loc.in_lru && !loc.is_valid) { + device(DeviceId(i)).deallocate(system(), stream_hint, buf); } - - if (!victim->host_entry.is_allocated) { - allocate_host(*victim, device_id); - } - - copy_d2h(device_id, *victim); } } - spdlog::debug( - "evict buffer {} from GPU {}, frees {} bytes", - victim->name, - device_id, - victim->layout.size_in_bytes - ); - - deallocate_device_async(device_id, *victim); - return true; -} - -AllocationResult MemoryManager::try_allocate_device_async(DeviceId device_id, Buffer& buffer) { - auto& device_entry = buffer.device_entry[device_id]; - - if (device_entry.is_allocated) { - return AllocationResult::Success; - } + refcnt_ptr memory_system; + uint64_t next_request_id_counter = 1; + uint64_t next_transaction_id_counter = 1; + HostState host_state; + DeviceState device_states[MAX_DEVICES]; +}; - KMM_ASSERT(device_entry.num_allocation_locks == 0); - spdlog::trace( - "allocate {} bytes on GPU {} for buffer {}", - buffer.layout.size_in_bytes, - device_id, - buffer.name - ); +MemoryManager::MemoryManager(refcnt_ptr memory_system) : + m_impl(std::make_unique(std::move(memory_system))) {} - g_device_ptr_t ptr_out; - DeviceEventSet events; - auto result = - m_memory->allocate_device(device_id, buffer.layout.size_in_bytes, &ptr_out, events); +MemoryManager::~MemoryManager() { + // By this point, every buffer must have gone through `release_buffer` (the caller's + // `RuntimeImpl` sweeps any it still owns before destroying `MemoryManager`), so there should + // be nothing left allocated and no request still waiting for memory. + KMM_ASSERT(m_impl->host().bytes_allocated == 0); - if (result != AllocationResult::Success) { - return result; + for (size_t id = 0; id < MAX_DEVICES; id++) { + auto& device = m_impl->device(DeviceId(id)); + KMM_ASSERT(device.bytes_allocated == 0); + KMM_ASSERT(device.req_head == nullptr); } +} - device_entry.data = ptr_out; - device_entry.is_allocated = true; - device_entry.is_valid = false; - device_entry.epoch_event = events; - device_entry.access_events = events; - device_entry.write_events = std::move(events); +MemoryBuffer MemoryManager::create_buffer( + std::unique_ptr data, + std::string name, + bool evictable, + std::optional home_memory_id +) { + spdlog::debug("buffer {} has been created", name); - device_at(device_id).add_to_lru(buffer); - check_consistency(); - return result; + return MemoryBuffer( + make_refcnt(std::move(name), evictable, std::move(data), home_memory_id) + ); } -void MemoryManager::deallocate_device_async(DeviceId device_id, Buffer& buffer) { - auto& device_entry = buffer.device_entry[device_id]; - KMM_ASSERT(device_entry.num_allocation_locks == 0); +void MemoryManager::release_buffer(MemoryBuffer buffer) { + auto* buf = buffer.get(); - if (!device_entry.is_allocated) { + // check if already released + if (buf->released) { return; } - size_t size_in_bytes = buffer.layout.size_in_bytes; - spdlog::trace( - "free {} bytes for buffer {} on GPU {} (dependencies={})", - size_in_bytes, - buffer.name, - device_id, - device_entry.access_events - ); - - m_memory->deallocate_device( - device_id, - device_entry.data, - size_in_bytes, - std::move(device_entry.access_events) - ); + buf->released = true; + KMM_ASSERT(buf->queue_head == nullptr); + auto stream_hint = DeviceStreamId::null(); - device_at(device_id).remove_from_lru(buffer); + m_impl->host().deallocate(m_impl->system(), stream_hint, buf); - device_entry.is_allocated = false; - device_entry.is_valid = false; - device_entry.epoch_event.clear(); - device_entry.write_events.clear(); - device_entry.access_events.clear(); - device_entry.data = 0; + for (size_t id = 0; id < MAX_DEVICES; id++) { + m_impl->device(DeviceId(id)).deallocate(m_impl->system(), stream_hint, buf); + } - check_consistency(); + spdlog::debug("buffer {} has been deleted", buffer->name); } -void MemoryManager::lock_allocation_host(Buffer& buffer, DeviceId device_affinity, Request& req) { - auto& host_entry = buffer.host_entry; +MemoryTransaction MemoryManager::create_transaction(MemoryTransaction parent) { + auto id = m_impl->next_transaction_id_counter++; + auto txn = MemoryTransaction(make_refcnt(id, std::move(parent))); - if (!host_entry.is_allocated) { - allocate_host(buffer, device_affinity); + if (txn->parent) { + spdlog::debug("transaction {} has been created (parent={})", id, txn->parent->id); + } else { + spdlog::debug("transaction {} has been created", id); } - host_entry.num_allocation_locks++; - spdlog::trace( - "lock allocation on host of buffer {} for request {}", - buffer.name, - req.identifier - ); + return txn; } -void MemoryManager::unlock_allocation_host(Buffer& buffer, Request& req) { - auto& host_entry = buffer.host_entry; +MemoryRequest MemoryManager::create_request( + const MemoryBuffer& buffer, + MemoryId memory_id, + AccessKind mode, + MemoryTransaction parent, + NotifyHandle callback +) { + auto req = make_refcnt( + m_impl->next_request_id_counter++, + refcnt_ptr(buffer.get(), true), + memory_id, + mode, + std::move(callback) + ); - KMM_ASSERT(host_entry.is_allocated); - KMM_ASSERT(host_entry.num_allocation_locks > 0); + // register with buffer + if (req->buffer->try_register_request(req.get())) { + // set state to waiting for allocation + req->state = MemoryRequestImpl::State::WaitingForAllocation; + req->parent = parent; - host_entry.num_allocation_locks--; - spdlog::trace( - "unlock allocation on host of buffer {} for request {}", - buffer.name, - req.identifier + // register with memory arbiter + if (req->memory_id.is_device()) { + m_impl->device(req->memory_id.as_device()).register_request(req.get()); + } + } else { + throw std::runtime_error("failed to lock buffer for access"); + } + + spdlog::debug( + "request {} has been created for buffer {} (memory={}, access={})", + req->id, + buffer->name, + memory_id, + mode ); + return req; } -bool MemoryManager::try_lock_allocation_device(DeviceId device_id, Buffer& buffer, Request& req) { - KMM_ASSERT(req.allocation_acquired == false); +Poll MemoryManager::poll_request(const DeviceStreamId& stream_hint, const MemoryRequest& request) { + MemoryManager::Impl& mgr = *m_impl; + MemoryRequestImpl* req = request.get(); + auto* buf = req->buffer.get(); + auto memory_id = req->memory_id; - if (device_at(device_id).allocation_first_pending != &req) { - return false; + if (req->state == MemoryRequestImpl::State::Unqueued) { + throw std::runtime_error("cannot poll memory request not registered with transaction"); } - while (true) { - // Try to allocate - auto result = try_allocate_device_async(device_id, buffer); - - if (result == AllocationResult::Success) { - break; + if (req->state == MemoryRequestImpl::State::WaitingForAllocation) { + if (memory_id.is_host()) { + mgr.host().acquire_for_request(mgr.system(), stream_hint, buf); + } else { + if (!mgr.device(memory_id.as_device()) + .try_acquire_for_request(mgr.system(), stream_hint, req)) { + return Poll::Pending; + } } - // No memory available, try to free memory - if (try_free_device_memory(device_id)) { - continue; - } + spdlog::debug("request {} has been granted memory for buffer {}", req->id, buf->name); + req->state = MemoryRequestImpl::State::Granted; + } - if (!is_out_of_memory(device_id, req)) { - return false; + if (req->state == MemoryRequestImpl::State::Granted) { + if (buf->before_access(mgr.system(), stream_hint, memory_id, req->mode) == Poll::Pending) { + return Poll::Pending; } - throw std::runtime_error(fmt::format( - "cannot allocate {} bytes on GPU {}, out of memory", - buffer.layout.size_in_bytes, - device_id - )); + mgr.reclaim_invalidated(stream_hint, buf); + req->state = MemoryRequestImpl::State::Ready; } - spdlog::trace( - "lock allocation on GPU {} of buffer {} for request {}", - device_id, - buffer.name, - req.identifier - ); - - device_at(device_id).increment_allocation_locks(buffer, req); - return true; -} - -void MemoryManager::unlock_allocation_device( - DeviceId device_id, - Buffer& buffer, - Request& req -) noexcept { - spdlog::trace( - "unlock allocation on GPU {} of buffer {} for request {}", - device_id, - buffer.name, - req.identifier - ); - - device_at(device_id).decrement_allocation_locks(buffer); - check_consistency(); + spdlog::debug("request {} has been granted access to buffer {}", req->id, buf->name); + KMM_ASSERT(req->state == MemoryRequestImpl::State::Ready); + return Poll::Ready; } -void MemoryManager::prepare_access_to_buffer( - MemoryId memory_id, - Buffer& buffer, - AccessMode mode, +BufferAccessor MemoryManager::access_request( + const MemoryRequest& request, DeviceEventSet& deps_out ) { - bool is_writer = mode != AccessMode::Read; - bool is_exclusive = mode == AccessMode::Exclusive; - auto& entry = buffer.entry(memory_id); + auto* req = request.get(); + auto* buf = req->buffer.get(); - make_entry_valid(memory_id, buffer, deps_out); - - if (is_writer) { - make_entry_exclusive(memory_id, buffer, deps_out); - } - - if (is_exclusive) { - deps_out.insert(entry.access_events); - } + KMM_ASSERT(req->state == MemoryRequestImpl::State::Ready); + return buf->access(req->memory_id, req->mode, deps_out); } -void MemoryManager::finalize_access_to_buffer( - MemoryId memory_id, - Buffer& buffer, - AccessMode mode, - DeviceEvent event -) noexcept { - bool is_writer = mode != AccessMode::Read; - bool is_exclusive = mode == AccessMode::Exclusive; - auto& entry = buffer.entry(memory_id); - - entry.access_events.insert(event); - - if (is_writer) { - entry.write_events.insert(event); - } +void MemoryManager::release_request(MemoryRequest request, const DeviceEventSet& deps) { + MemoryManager::Impl& mgr = *m_impl; + auto* req = request.get(); + auto* buf = request->buffer.get(); + auto memory_id = request->memory_id; - if (is_exclusive) { - entry.epoch_event.insert(event); + if (req->state == MemoryRequestImpl::State::Ready) { + buf->after_access(memory_id, req->mode, deps); + req->state = MemoryRequestImpl::State::Granted; } -} -std::optional MemoryManager::find_valid_device_entry(const Buffer& buffer) { - for (size_t device_id = 0; device_id < MAX_DEVICES; device_id++) { - if (buffer.device_entry[device_id].is_valid) { - return DeviceId(device_id); + if (req->state == MemoryRequestImpl::State::Granted) { + if (memory_id.is_device()) { + mgr.device(memory_id.as_device()).release_for_request(buf); + } else { + mgr.host().release_for_request(buf); } - } - - return std::nullopt; -} - -void MemoryManager::make_entry_valid(MemoryId memory_id, Buffer& buffer, DeviceEventSet& deps_out) { - auto& entry = buffer.entry(memory_id); - - KMM_ASSERT(entry.is_allocated); - deps_out.insert(entry.epoch_event); - if (entry.is_valid) { - return; + req->state = MemoryRequestImpl::State::WaitingForAllocation; } - if (memory_id.is_host()) { - if (auto src_id = find_valid_device_entry(buffer)) { - deps_out.insert(copy_d2h(*src_id, buffer)); - return; - } - } else { - auto device_id = memory_id.as_device(); - - if (buffer.host_entry.is_valid) { - deps_out.insert(copy_h2d(device_id, buffer)); - return; + if (req->state == MemoryRequestImpl::State::WaitingForAllocation) { + if (memory_id.is_device()) { + mgr.device(memory_id.as_device()).unregister_request(req); } - if (auto src_id = find_valid_device_entry(buffer)) { - if (m_memory->is_copy_supported(*src_id, memory_id)) { - deps_out.insert(copy_d2d(*src_id, device_id, buffer)); - } else { - if (!buffer.host_entry.is_allocated) { - allocate_host(buffer, *src_id); - } - deps_out.insert(copy_d2h(*src_id, buffer)); - deps_out.insert(copy_h2d(device_id, buffer)); - } + buf->unregister_request(req); + req->state = MemoryRequestImpl::State::Unqueued; + } - return; - } + if (req->state == MemoryRequestImpl::State::Unqueued) { + // Never activated: not present in any queue, nothing to unwind. + req->state = MemoryRequestImpl::State::Released; } - entry.is_valid = true; + spdlog::debug("request {} released access to buffer {}", req->id, buf->name); + KMM_ASSERT(req->state == MemoryRequestImpl::State::Released); } -void MemoryManager::make_entry_exclusive( +void MemoryManager::prefetch_buffer( + const MemoryBuffer& buffer, MemoryId memory_id, - Buffer& buffer, - DeviceEventSet& deps_out + AccessKind mode ) { - make_entry_valid(memory_id, buffer, deps_out); + auto stream_hint = DeviceStreamId::null(); + MemoryRequest req; - // invalidate host if necessary - if (memory_id != MemoryId::host()) { - buffer.host_entry.is_valid = false; - deps_out.insert(buffer.host_entry.access_events); - } + // A single poll attempt: this is only a hint, so if the buffer cannot be + // allocated/granted right away (contended, or out of memory), give up + // instead of waiting like a real access would. `release_request` cleanly + // unwinds whatever partial progress was made (e.g. allocated but not yet + // granted), same as it does for any request that never reaches `Ready`. + DeviceEventSet deps; - // Invalidate all _other_ device entries - for (size_t i = 0; i < MAX_DEVICES; i++) { - if (memory_id == MemoryId(DeviceId(i))) { - continue; - } + try { + req = create_request(buffer, memory_id, mode); - auto& peer_entry = buffer.device_entry[i]; - peer_entry.is_valid = false; - deps_out.insert(peer_entry.access_events); + if (poll_request(stream_hint, req) == Poll::Ready) { + // Say hi! + access_request(req, deps); + } + } catch (const std::exception&) { + // e.g. out of memory, drop the hint. } -} - -DeviceEvent MemoryManager::copy_h2d(DeviceId device_id, Buffer& buffer) { - spdlog::trace( - "copy {} bytes from host to GPU {} for buffer {}", - buffer.layout.size_in_bytes, - device_id, - buffer.name - ); - - auto& host_entry = buffer.host_entry; - auto& device_entry = buffer.device_entry[device_id]; - - KMM_ASSERT(host_entry.is_allocated && device_entry.is_allocated); - KMM_ASSERT(host_entry.is_valid && !device_entry.is_valid); - - DeviceEventSet deps = device_entry.access_events | host_entry.write_events; - auto event = m_memory->copy_host_to_device( - device_id, - host_entry.data, - device_entry.data, - buffer.layout.size_in_bytes, - std::move(deps) - ); - - host_entry.access_events.insert(event); - device_entry.epoch_event = {event}; - device_entry.access_events = {event}; - device_entry.write_events = {event}; - - device_entry.is_valid = true; - return event; -} - -DeviceEvent MemoryManager::copy_d2h(DeviceId device_id, Buffer& buffer) { - spdlog::trace( - "copy {} bytes from GPU {} to host for buffer {}", - buffer.layout.size_in_bytes, - device_id, - buffer.name - ); - - auto& host_entry = buffer.host_entry; - auto& device_entry = buffer.device_entry[device_id]; - - KMM_ASSERT(host_entry.is_allocated && device_entry.is_allocated); - KMM_ASSERT(!host_entry.is_valid && device_entry.is_valid); - - DeviceEventSet deps = device_entry.write_events | host_entry.access_events; - - auto event = m_memory->copy_device_to_host( - device_id, - device_entry.data, - host_entry.data, - buffer.layout.size_in_bytes, - std::move(deps) - ); - - device_entry.access_events.insert(event); - host_entry.epoch_event.insert(event); - host_entry.access_events.insert(event); - host_entry.write_events.insert(event); - host_entry.is_valid = true; - - return event; + if (req != nullptr) { + release_request(req, deps); + } } -DeviceEvent MemoryManager::copy_d2d( - DeviceId device_src_id, - DeviceId device_dst_id, - Buffer& buffer -) { - spdlog::trace( - "copy {} bytes from GPU {} to GPU {} for buffer {}", - buffer.layout.size_in_bytes, - device_src_id, - device_dst_id, - buffer.name - ); +void MemoryManager::try_evict_buffer(const MemoryBuffer& buffer, MemoryId memory_id) { + auto stream_hint = DeviceStreamId::null(); - auto& src_entry = buffer.device_entry[device_src_id]; - auto& dst_entry = buffer.device_entry[device_dst_id]; + if (!buffer->is_allocated(memory_id)) { + return; + } - KMM_ASSERT(src_entry.is_allocated && dst_entry.is_allocated); - KMM_ASSERT(src_entry.is_valid && !dst_entry.is_valid); + if (memory_id.is_device()) { + m_impl->device(memory_id.as_device()) + .try_evict_buffer(m_impl->system(), stream_hint, buffer.get()); + return; + } - DeviceEventSet deps = dst_entry.access_events | src_entry.write_events; + // Host memory has no further fallback location: unlike a device eviction, + // which can always fall back to copying into host memory first, deallocating + // the host location while it holds the buffer's only valid copy would destroy + // the data outright. Bail instead of doing that. + auto& loc = buffer->host_location; - auto event = m_memory->copy_device_to_device( - device_src_id, - device_dst_id, - src_entry.data, - dst_entry.data, - buffer.layout.size_in_bytes, - std::move(deps) - ); + if (loc.alloc_count > 0) { + return; + } - src_entry.access_events.insert(event); - dst_entry.epoch_event = {event}; - dst_entry.access_events = {event}; - dst_entry.write_events = {event}; - dst_entry.is_valid = true; + MemoryId other = MemoryId::host(); + if (loc.is_valid && !buffer->find_valid_location(MemoryId::host(), other)) { + return; + } - return event; + m_impl->host_state.deallocate(m_impl->system(), stream_hint, buffer.get()); } -MemoryManager::Device& MemoryManager::device_at(DeviceId id) noexcept { - KMM_ASSERT(id < MAX_DEVICES); - return m_devices[id]; +void MemoryManager::invalidate_buffer(const MemoryBuffer& buffer) { + auto stream_hint = DeviceStreamId::null(); + spdlog::debug("buffer {} has been invalidated", buffer->name); + buffer->invalidate_all(); + m_impl->reclaim_invalidated(stream_hint, buffer.get()); } -bool MemoryManager::is_out_of_memory(DeviceId device_id, Request& req) { - std::unordered_set waiting_transactions; - auto& device = device_at(device_id); - - // First, iterate over the requests that are waiting for allocation. Mark all the - // related transactions as `waiting` by adding them to `waiting_transactions` - for (auto* it = device.allocation_first_pending; it != nullptr; it = it->allocation_next) { - auto* p = it->parent.get(); - - while (p != nullptr) { - waiting_transactions.insert(p); - p = p->parent.get(); - } - } - - // Next, iterate over the requests that have been granted an allocation. If the associated - // transaction of one of the requests has not been marked as waiting, we are not out of memory - // since that transaction will release its memory again at some point in the future. - for (auto* it = device.allocation_head; it != device.allocation_first_pending; - it = it->allocation_next) { - auto* p = it->parent.get(); - - if (waiting_transactions.find(p) == waiting_transactions.end()) { - return false; +void MemoryManager::trim_device(DeviceId id, size_t bytes_remaining, bool evict) { + auto& device = m_impl->device(id); + auto& system = m_impl->system(); + auto stream_hint = DeviceStreamId::null(); + auto memory_id = MemoryId::device(id); + auto bytes_before = system.bytes_reserved(memory_id); + + // repeatedly try to evict a buffer until we get under the limit. + if (evict) { + // 1. attempt to get bytes_allocated below the given limit. + while (device.bytes_allocated > bytes_remaining) { + if (!device.try_evict_one(system, stream_hint)) { + break; + } } - } - - spdlog::error( - "out of memory for GPU {}, failed to allocate {} bytes for request {} of buffer {}", - device_id, - req.buffer->layout.size_in_bytes, - req.buffer->name, - req.identifier - ); - - spdlog::error("following buffers are currently allocated: "); - for (const auto& buffer : m_buffers) { - auto& entry = buffer->entry(device_id); + // 2. attempt to get bytes_reserved below the given limit. + while (system.bytes_reserved(memory_id) > bytes_remaining) { + if (!device.try_evict_one(system, stream_hint)) { + break; + } - if (entry.is_allocated) { - spdlog::error( - " - buffer {} ({} bytes, {} requests active, {} allocation locks)", - buffer->name, - buffer->layout.size_in_bytes, - buffer->num_requests_active, - entry.num_allocation_locks - ); + // if still not below the limit, try to trim memory and see if that helps. + if (system.bytes_reserved(memory_id) > bytes_remaining) { + system.trim_device(id, bytes_remaining); + } } } - spdlog::error("following requests have been granted:"); - for (auto* it = device.allocation_head; it != device.allocation_first_pending; - it = it->allocation_next) { - spdlog::error( - " - request {} ({} bytes, buffer {}, transaction {})", - it->identifier, - it->buffer->layout.size_in_bytes, - it->buffer->name, - it->parent->id - ); - } + // trim memory + system.trim_device(id, bytes_remaining); - spdlog::error("following requests are pending:"); - for (auto* it = device.allocation_first_pending; it != nullptr; it = it->allocation_next) { - spdlog::error( - " - request {} ({} bytes, buffer {}, transaction {})", - it->identifier, - it->buffer->layout.size_in_bytes, - it->buffer->name, - it->parent->id + auto bytes_after = system.bytes_reserved(memory_id); + if (bytes_after != bytes_before) { + spdlog::debug( + "trimmed device {} from {} to {} bytes reserved (target: {})", + id, + bytes_before, + bytes_after, + bytes_remaining ); } - - return true; } -void MemoryManager::check_consistency() const {} - -// This is here to check the consistency of the data structures while debugging. -/* -void MemoryManager::check_consistency() const { - for (size_t i = 0; i < MAX_DEVICES; i++) { - std::unordered_set available_buffers; - auto id = DeviceId(i); - auto& device = m_devices[i]; - - auto* prev = (Buffer*) nullptr; - auto* current = device.lru_oldest; - - while (current != nullptr) { - available_buffers.insert(current); - auto& entry = current->device_entry[id]; - - KMM_ASSERT(entry.num_allocation_locks == 0); - KMM_ASSERT(entry.lru_older == prev); - - prev = current; - current = entry.lru_newer; - } - - KMM_ASSERT(prev == device.lru_newest); - - for (const auto& buffer: m_buffers) { - auto& entry = buffer->device_entry[id]; +void MemoryManager::make_progress() { + // maybe in the future... +} - if (entry.is_allocated && entry.num_allocation_locks == 0) { - KMM_ASSERT(available_buffers.find(buffer.get()) != available_buffers.end()); - } - } +std::ostream& operator<<(std::ostream& stream, AccessKind access) { + switch (access) { + case AccessKind::ReadOnly: + return stream << "ReadOnly"; + case AccessKind::SharedWrite: + return stream << "SharedWrite"; + case AccessKind::Exclusive: + return stream << "Exclusive"; + default: + return stream << "???"; } } -*/ + +KMM_REFCNT_TRAITS_IMPL(MemoryTransactionImpl) +KMM_REFCNT_TRAITS_IMPL(MemoryBufferImpl) +KMM_REFCNT_TRAITS_IMPL(MemoryRequestImpl) } // namespace kmm diff --git a/src/runtime/memory_system.cpp b/src/runtime/memory_system.cpp index 59bd5a00..99cc76fe 100644 --- a/src/runtime/memory_system.cpp +++ b/src/runtime/memory_system.cpp @@ -1,211 +1,876 @@ -#include +#include +#include #include "spdlog/spdlog.h" -#include "kmm/memops/gpu_fill.hpp" -#include "kmm/memops/host_fill.hpp" +#include "kmm/core/checked_math.hpp" +#include "kmm/runtime/allocators/arena.hpp" +#include "kmm/runtime/allocators/device.hpp" +#include "kmm/runtime/allocators/device_pool.hpp" +#include "kmm/runtime/allocators/limit.hpp" +#include "kmm/runtime/allocators/managed.hpp" +#include "kmm/runtime/allocators/pinned.hpp" +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/runtime/memops/fill_gpu.hpp" +#include "kmm/runtime/memops/reduction_gpu.hpp" #include "kmm/runtime/memory_system.hpp" +#include "kmm/utils/gpu_utils.hpp" namespace kmm { -struct MemorySystemImpl::Device { - KMM_NOT_COPYABLE(Device) +struct MemorySystem::DeviceState { + DeviceState(g_context_t context, g_device_t ordinal, std::unique_ptr allocator) : + context(context), + ordinal(ordinal), + allocator(std::move(allocator)) {} - public: - GPUContextHandle context; - std::unique_ptr allocator; + g_context_t context; + g_device_t ordinal; + std::unique_ptr allocator; - DeviceStream h2d_stream; - DeviceStream d2h_stream; - DeviceStream h2d_hi_stream; // high priority stream - DeviceStream d2h_hi_stream; // high priority stream + MemoryStats stats; - Device( - GPUContextHandle context, - std::unique_ptr allocator, - DeviceStreamManager& streams - ) : - context(context), - allocator(std::move(allocator)), - h2d_stream(streams.create_stream(context, false)), - d2h_stream(streams.create_stream(context, false)), - h2d_hi_stream(streams.create_stream(context, true)), - d2h_hi_stream(streams.create_stream(context, true)) {} + void record_allocation(size_t nbytes) { + stats.record_allocation(nbytes); + } + + void record_deallocation(size_t nbytes) { + stats.record_deallocation(nbytes); + } }; -MemorySystemImpl::MemorySystemImpl( - std::shared_ptr stream_manager, - std::vector device_contexts, - std::unique_ptr host_mem, - std::vector> device_mems -) : - m_streams(stream_manager), - m_host(std::move(host_mem)) - -{ - KMM_ASSERT(device_contexts.size() == device_mems.size()); - KMM_ASSERT(device_contexts.size() <= MAX_DEVICES); - - for (size_t i = 0; i < device_contexts.size(); i++) { - m_devices[i] = std::make_unique( - device_contexts[i], - std::move(device_mems[i]), - *stream_manager +static std::unique_ptr make_host_allocator( + const RuntimeConfig& config, + DeviceEventRegistry events, + g_context_t context +) { + std::unique_ptr allocator; + allocator = std::make_unique(context); + + if (config.host_memory_limit != std::numeric_limits::max()) { + allocator = std::make_unique( + std::move(allocator), + events, + config.host_memory_limit ); } + + if (config.host_memory_kind != HostMemoryKind::NoPool && config.host_memory_block_size > 0) { + allocator = + std::make_unique(std::move(allocator), config.host_memory_block_size); + } + + return allocator; } -MemorySystemImpl::~MemorySystemImpl() {} +static std::unique_ptr make_device_allocator( + const RuntimeConfig& config, + DeviceEventRegistry events, + g_context_t context, + size_t device_memory_size +) { + std::unique_ptr allocator; + + size_t limit = config.device_memory_limit; + + if (config.device_memory_keep_free > 0) { + if (device_memory_size <= config.device_memory_keep_free) { + throw std::runtime_error( + fmt::format( + "cannot reserve {} bytes on GPU, only {} bytes are available", + config.device_memory_keep_free, + device_memory_size + ) + ); + } -void MemorySystemImpl::make_progress() { - m_host->make_progress(); + limit = std::min(limit, device_memory_size - config.device_memory_keep_free); + } - for (const auto& device : m_devices) { - if (device == nullptr) { + switch (config.device_memory_kind) { + case DeviceMemoryKind::DefaultPool: + allocator = + std::make_unique(context, DevicePoolKind::Default, limit); + break; + case DeviceMemoryKind::PrivatePool: + allocator = + std::make_unique(context, DevicePoolKind::Create, limit); + break; + default: + allocator = std::make_unique(context); break; + } + + if (limit != std::numeric_limits::max()) { + allocator = std::make_unique(std::move(allocator), events, limit); + } + + if (config.device_memory_kind != DeviceMemoryKind::NoPool + && config.device_memory_block_size > 0) { + allocator = + std::make_unique(std::move(allocator), config.device_memory_block_size); + } + + return allocator; +} + +MemorySystem::MemorySystem( + const SystemInfo& system_info, + DeviceEventRegistry events, + const RuntimeConfig& config +) : + m_events(events), + m_streams( + system_info, + events, + config.device_concurrent_streams, + config.device_concurrent_streams, + config.device_concurrent_streams + ), + m_num_devices(system_info.num_devices()) { + spdlog::info("initializing memory system with {} device(s)", m_num_devices); + + g_context_t host_context = + m_num_devices > 0 ? system_info.device(DeviceId(0)).context() : nullptr; + m_host_allocator = make_host_allocator(config, events, host_context); + + if (m_num_devices > 0) { + m_managed_allocator = std::make_unique(host_context); + } + + for (size_t i = 0; i < m_num_devices; i++) { + const auto& info = system_info.device(DeviceId(i)); + + m_devices[i] = std::make_unique( + info.context(), + info.device_ordinal(), + make_device_allocator(config, events, info.context(), info.total_memory_size()) + ); + } + + // Determine, and where possible enable, peer-to-peer access between every device pair. + for (size_t i = 0; i < m_num_devices; i++) { + m_peer_access[i][i] = true; + + for (size_t j = i + 1; j < m_num_devices; j++) { + int i_can_access_j = 0; + int j_can_access_i = 0; + + KMM_GPU_CHECK(g_device_can_access_peer( + &i_can_access_j, + system_info.device(DeviceId(i)).device_ordinal(), + system_info.device(DeviceId(j)).device_ordinal() + )); + + KMM_GPU_CHECK(g_device_can_access_peer( + &j_can_access_i, + system_info.device(DeviceId(j)).device_ordinal(), + system_info.device(DeviceId(i)).device_ordinal() + )); + + m_peer_access[i][j] = i_can_access_j != 0; + m_peer_access[j][i] = j_can_access_i != 0; + + if (i_can_access_j) { + GPUContextGuard guard {m_devices[i]->context}; + g_result_t result = g_ctx_enable_peer_access(m_devices[j]->context, 0); + + if (result != GPU_ERROR_PEER_ACCESS_ALREADY_ENABLED) { + KMM_GPU_CHECK(result); + } + } + + if (j_can_access_i) { + GPUContextGuard guard {m_devices[j]->context}; + g_result_t result = g_ctx_enable_peer_access(m_devices[i]->context, 0); + + if (result != GPU_ERROR_PEER_ACCESS_ALREADY_ENABLED) { + KMM_GPU_CHECK(result); + } + } + + spdlog::debug( + "peer access between device {} and device {}: {} -> {} = {}, {} -> {} = {}", + i, + j, + i, + j, + i_can_access_j != 0, + j, + i, + j_can_access_i != 0 + ); } + } +} + +MemorySystem::~MemorySystem() { + spdlog::info("memory system stats:"); + spdlog::info(" - host:"); + spdlog::info(" - allocated: {} bytes", m_host_stats.bytes_allocated); + spdlog::info(" - allocated at peak: {} bytes", m_host_stats.max_bytes_inuse); + + for (size_t i = 0; i < m_num_devices; i++) { + spdlog::info(" - copied to device {}: {} bytes", i, m_host_stats.bytes_to_device[i]); + spdlog::info(" - copied from device {}: {} bytes", i, m_devices[i]->stats.bytes_to_host); + } - device->allocator->make_progress(); + if (auto reserved = m_host_allocator->bytes_reserved()) { + spdlog::info(" - reserved from underlying allocator: {} bytes", *reserved); + } + + for (size_t i = 0; i < m_num_devices; i++) { + const auto& state = *m_devices[i]; + + spdlog::info(" - device {}:", i); + spdlog::info(" - allocated: {} bytes", state.stats.bytes_allocated); + spdlog::info(" - allocated at peak: {} bytes", state.stats.max_bytes_inuse); + spdlog::info(" - copied to host: {} bytes", state.stats.bytes_to_host); + spdlog::info(" - copied from host: {} bytes", m_host_stats.bytes_to_device[i]); + + for (size_t j = 0; j < m_num_devices; j++) { + if (j != i) { + spdlog::info( + " - copied to device {}: {} bytes", + j, + state.stats.bytes_to_device[j] + ); + spdlog::info( + " - copied from device {}: {} bytes", + j, + m_devices[j]->stats.bytes_to_device[i] + ); + } + } + + if (auto reserved = state.allocator->bytes_reserved()) { + spdlog::info(" - reserved from underlying allocator: {} bytes", *reserved); + } } } -void MemorySystemImpl::trim_host(size_t bytes_remaining) { - m_host->trim(bytes_remaining); +MemorySystem::DeviceState& MemorySystem::device_state(DeviceId id) const { + KMM_ASSERT(id.get() < m_num_devices); + return *m_devices[id.get()]; } -AllocationResult MemorySystemImpl::allocate_host( - size_t nbytes, - DeviceId device_affinity, +void MemorySystem::make_progress() { + m_host_allocator->poll(); + + for (size_t i = 0; i < m_num_devices; i++) { + m_devices[i]->allocator->poll(); + } + + m_streams.make_progress(); +} + +void MemorySystem::trim_host(size_t bytes_remaining) { + m_host_allocator->trim(bytes_remaining); +} + +void MemorySystem::trim_device(DeviceId id, size_t bytes_remaining) { + device_state(id).allocator->trim(bytes_remaining); +} + +size_t MemorySystem::bytes_reserved(MemoryId id) const { + if (id.is_host()) { + return m_host_allocator->bytes_reserved().value_or(m_host_stats.bytes_inuse); + } + + const auto& state = device_state(id.as_device()); + return state.allocator->bytes_reserved().value_or(state.stats.bytes_inuse); +} + +template +DeviceEvent schedule_onto_stream( + bool use_hint, + const DeviceEventRegistry& events, + DeviceDataStreams& streams, + DeviceId device_id, + StreamKind kind, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + F callback +) { + if (stream_hint.is_null() || !use_hint) { + return streams.submit(device_id, kind, deps_in, [&](auto stream_id) -> uint64_t { + return callback(DeviceStream {events, stream_id}); + }); + } else { + return events.submit(stream_hint, deps_in, [&](auto stream) { + callback(DeviceStream {events, stream_hint}); + }); + } +} + +DeviceId MemorySystem::affinity_for_stream(const DeviceStreamId& stream_hint) { + if (!stream_hint.is_null()) { + for (size_t i = 0; i < m_num_devices; i++) { + if (m_events.has_context(stream_hint, m_devices[i]->context)) { + return DeviceId(i); + } + } + } + + return DeviceId(0); +} + +AllocResult MemorySystem::allocate_host( + BufferLayout layout, void** ptr_out, + const DeviceStreamId& stream_hint, DeviceEventSet& deps_out ) { - // TODO: take into account device_affinity + AllocResult result; + + if (stream_hint.is_null()) { + result = m_host_allocator->allocate(layout, ptr_out); + } else { + result = + m_host_allocator->allocate_async(DeviceStream {m_events, stream_hint}, layout, ptr_out); - auto result = m_host->allocate_async(nbytes, ptr_out, deps_out); - if (result != AllocationResult::Success) { - return result; + deps_out.insert(m_events.record(stream_hint)); } - deps_out.remove_ready(*m_streams); - return AllocationResult::Success; + if (result == AllocResult::Success) { + m_host_stats.record_allocation(layout.size_in_bytes); + } + + spdlog::trace( + "allocate {} bytes of host memory (addr: {}, result: {})", + layout.size_in_bytes, + *ptr_out, + static_cast(result) + ); + + return result; } -void MemorySystemImpl::deallocate_host(void* ptr, size_t nbytes, DeviceEventSet deps) { - deps.remove_ready(*m_streams); - m_host->deallocate_async(ptr, nbytes, std::move(deps)); +void MemorySystem::deallocate_host( + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in +) { + spdlog::trace("deallocate {} bytes of host memory (addr: {})", layout.size_in_bytes, ptr); + + m_host_stats.record_deallocation(layout.size_in_bytes); + + if (stream_hint.is_null()) { + m_host_allocator->deallocate(ptr, layout); + } else { + m_events.wait_on_event(stream_hint, deps_in); + m_host_allocator->deallocate_async(DeviceStream {m_events, stream_hint}, ptr, layout); + } } -void MemorySystemImpl::trim_device(size_t bytes_remaining) { - for (const auto& device : m_devices) { - if (device != nullptr) { - device->allocator->trim(bytes_remaining); - } +AllocResult MemorySystem::allocate_managed( + BufferLayout layout, + void** ptr_out, + const DeviceStreamId& stream_hint, + DeviceEventSet& deps_out +) { + if (!m_managed_allocator) { + return AllocResult::ErrorUnsupported; + } + + AllocResult result; + + if (stream_hint.is_null()) { + result = m_managed_allocator->allocate(layout, ptr_out); + } else { + result = m_managed_allocator->allocate_async( + DeviceStream {m_events, stream_hint}, + layout, + ptr_out + ); + + deps_out.insert(m_events.record(stream_hint)); + } + + if (result == AllocResult::Success) { + m_managed_stats.record_allocation(layout.size_in_bytes); + } + + spdlog::trace( + "allocate {} bytes of managed memory (addr: {}, result: {})", + layout.size_in_bytes, + *ptr_out, + static_cast(result) + ); + + return result; +} + +void MemorySystem::deallocate_managed( + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in +) { + spdlog::trace("deallocate {} bytes of managed memory (addr: {})", layout.size_in_bytes, ptr); + + m_managed_stats.record_deallocation(layout.size_in_bytes); + + if (stream_hint.is_null()) { + m_managed_allocator->deallocate(ptr, layout); + } else { + m_events.wait_on_event(stream_hint, deps_in); + m_managed_allocator->deallocate_async(DeviceStream {m_events, stream_hint}, ptr, layout); + } +} + +void MemorySystem::prefetch_managed( + MemoryId memory_id, + void* ptr, + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps +) { + if (stream_hint.is_null() || !memory_id.is_device()) { + return; + } + + auto device_id = memory_id.as_device(); + if (!same_context(device_id, stream_hint)) { + return; } + + spdlog::trace( + "prefetch {} bytes of managed memory to {} (addr: {})", + layout.size_in_bytes, + memory_id, + ptr + ); + + GPUContextGuard guard {device_state(device_id).context}; + m_events.wait_on_event(stream_hint, deps); + + KMM_GPU_CHECK(g_mem_prefetch_async( + reinterpret_cast(ptr), + layout.size_in_bytes, + device_state(device_id).ordinal, + m_events.get(stream_hint) + )); +} + +void* MemorySystem::translate_host_pointer(DeviceId device_id, void* host_ptr) const { + GPUContextGuard guard {device_state(device_id).context}; + + g_device_ptr_t dptr; + KMM_GPU_CHECK(g_mem_host_get_device_pointer(&dptr, host_ptr, 0)); + return reinterpret_cast(dptr); } -AllocationResult MemorySystemImpl::allocate_device( +bool MemorySystem::same_context(kmm::DeviceId device_id, const kmm::DeviceStreamId& stream_hint) { + return m_events.has_context(stream_hint, device_state(device_id).context); +} + +AllocResult MemorySystem::allocate_device( DeviceId device_id, - size_t nbytes, + BufferLayout layout, g_device_ptr_t* ptr_out, + const DeviceStreamId& stream_hint, DeviceEventSet& deps_out ) { - KMM_ASSERT(m_devices[device_id]); - auto& device = *m_devices[device_id]; - void* addr; + void* addr = nullptr; + AllocResult result; + + auto event = schedule_onto_stream( + same_context(device_id, stream_hint), + m_events, + m_streams, + device_id, + StreamKind::DeviceToDevice, + stream_hint, + {}, + [&](const auto& stream) { + result = device_state(device_id).allocator->allocate_async(stream, layout, &addr); + return 0; + } + ); - GPUContextGuard guard {device.context}; + deps_out.insert(event); + *ptr_out = reinterpret_cast(addr); - auto result = device.allocator->allocate_async(nbytes, &addr, deps_out); - if (result != AllocationResult::Success) { - return result; + if (result == AllocResult::Success) { + device_state(device_id).record_allocation(layout.size_in_bytes); } - deps_out.remove_ready(*m_streams); - *ptr_out = (g_device_ptr_t)addr; - return AllocationResult::Success; + spdlog::trace( + "allocate {} bytes of device memory on device {} (addr: {:#x}, result: {})", + layout.size_in_bytes, + device_id, + *ptr_out, + static_cast(result) + ); + + return result; } -void MemorySystemImpl::deallocate_device( +void MemorySystem::deallocate_device( DeviceId device_id, g_device_ptr_t ptr, - size_t nbytes, - DeviceEventSet deps + BufferLayout layout, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in ) { - deps.remove_ready(*m_streams); - - KMM_ASSERT(m_devices[device_id]); - auto& device = *m_devices[device_id]; + spdlog::trace( + "deallocate {} bytes of device memory on device {} (addr: {:#x})", + layout.size_in_bytes, + device_id, + ptr + ); + + schedule_onto_stream( + same_context(device_id, stream_hint), + m_events, + m_streams, + device_id, + StreamKind::DeviceToDevice, + stream_hint, + deps_in, + [&](const auto& stream) { + device_state(device_id).allocator->deallocate_async( + stream, + reinterpret_cast(ptr), + layout + ); + + return 0; + } + ); - GPUContextGuard guard {device.context}; - device.allocator->deallocate_async((void*)ptr, nbytes, std::move(deps)); + device_state(device_id).record_deallocation(layout.size_in_bytes); } -// Copies smaller than this threshold are put onto a high priority stream. This can improve -// performance since small copy jobs (like copying a single number) are prioritized over large -// slow copy jobs of several gigabytes. -static constexpr size_t HIGH_PRIORITY_THRESHOLD = 1024L * 1024; - -DeviceEvent MemorySystemImpl::copy_host_to_device( +DeviceEvent MemorySystem::copy_host_to_device( DeviceId device_id, const void* src_addr, g_device_ptr_t dst_addr, size_t nbytes, - DeviceEventSet deps + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in ) { - KMM_ASSERT(m_devices[device_id]); - auto& device = *m_devices[device_id]; - auto stream = nbytes <= HIGH_PRIORITY_THRESHOLD ? device.h2d_hi_stream : device.h2d_stream; + spdlog::trace( + "copy {} bytes from host (addr: {}) to device {} (addr: {:#x})", + nbytes, + src_addr, + device_id, + dst_addr + ); + + bool use_hint = + same_context(device_id, stream_hint) && m_events.is_latest_in(stream_hint, deps_in); + + auto event = schedule_onto_stream( + use_hint, + m_events, + m_streams, + device_id, + StreamKind::HostToDevice, + stream_hint, + deps_in, + [&](g_stream_t stream) { + KMM_GPU_CHECK(g_memcpy_h_to_d_async(dst_addr, src_addr, nbytes, (g_stream_t)stream)); + return nbytes; + } + ); - GPUContextGuard guard {device.context}; - return m_streams->with_stream(stream, deps, [&](auto stream) { - KMM_GPU_CHECK(g_memcpy_h_to_d_async(dst_addr, src_addr, nbytes, stream)); - }); + m_host_stats.bytes_to_device[device_id.get()] += nbytes; + return event; } -DeviceEvent MemorySystemImpl::copy_device_to_host( +DeviceEvent MemorySystem::copy_device_to_host( DeviceId device_id, g_device_ptr_t src_addr, void* dst_addr, size_t nbytes, - DeviceEventSet deps + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in ) { - KMM_ASSERT(m_devices[device_id]); - auto& device = *m_devices[device_id]; - auto stream = nbytes <= HIGH_PRIORITY_THRESHOLD ? device.d2h_hi_stream : device.d2h_stream; + spdlog::trace( + "copy {} bytes from device {} (addr: {:#x}) to host (addr: {})", + nbytes, + device_id, + src_addr, + dst_addr + ); + + bool use_hint = + same_context(device_id, stream_hint) && m_events.is_latest_in(stream_hint, deps_in); + + auto event = schedule_onto_stream( + use_hint, + m_events, + m_streams, + device_id, + StreamKind::DeviceToHost, + stream_hint, + deps_in, + [&](g_stream_t stream) { + KMM_GPU_CHECK(g_memcpy_d_to_h_async(dst_addr, src_addr, nbytes, (g_stream_t)stream)); + return nbytes; + } + ); - GPUContextGuard guard {device.context}; - return m_streams->with_stream(stream, deps, [&](auto stream) { - KMM_GPU_CHECK(g_memcpy_d_to_h_async(dst_addr, src_addr, nbytes, stream)); - }); + device_state(device_id).stats.bytes_to_host += nbytes; + return event; } -DeviceEvent MemorySystemImpl::copy_device_to_device( - DeviceId src_device_id, - DeviceId dst_device_id, +DeviceEvent MemorySystem::copy_device_to_device( + DeviceId src_device, + DeviceId dst_device, g_device_ptr_t src_addr, g_device_ptr_t dst_addr, size_t nbytes, - DeviceEventSet deps -) { - KMM_ASSERT(m_devices[dst_device_id] && m_devices[src_device_id]); - auto& src_device = *m_devices[src_device_id]; - auto& dst_device = *m_devices[dst_device_id]; - auto stream = - nbytes <= HIGH_PRIORITY_THRESHOLD ? dst_device.h2d_hi_stream : dst_device.h2d_stream; - - GPUContextGuard guard {dst_device.context}; - return m_streams->with_stream(stream, deps, [&](auto stream) { - KMM_GPU_CHECK(g_memcpy_peer_async( - dst_addr, - dst_device.context, - dst_device_id, - src_addr, - src_device.context, - src_device_id, - nbytes, - stream - )); + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in +) { + spdlog::trace( + "copy {} bytes from device {} (addr: {:#x}) to device {} (addr: {:#x})", + nbytes, + src_device, + src_addr, + dst_device, + dst_addr + ); + + auto event = m_streams.submit( // + dst_device, + StreamKind::DeviceToDevice, + deps_in, + [&](auto stream_id) -> uint64_t { + KMM_GPU_CHECK(g_memcpy_peer_async( + dst_addr, + device_state(dst_device).context, + device_state(dst_device).ordinal, + src_addr, + device_state(src_device).context, + device_state(src_device).ordinal, + nbytes, + m_events.get(stream_id) + )); + + return nbytes; + } + ); + + device_state(src_device).stats.bytes_to_device[dst_device.get()] += nbytes; + return event; +} + +AllocResult MemorySystem::allocate_host_and_copy_from_device( + BufferLayout layout, + void** dst_addr, + DeviceId device_id, + g_device_ptr_t src_addr, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEvent& dep_out +) { + AllocResult result; + bool use_hint = + same_context(device_id, stream_hint) && m_events.is_latest_in(stream_hint, deps_in); + + dep_out = schedule_onto_stream( + use_hint, + m_events, + m_streams, + device_id, + StreamKind::DeviceToHost, + stream_hint, + deps_in, + [&](const auto& stream) { + void* addr = nullptr; + + result = m_host_allocator->allocate_async(stream, layout, &addr); + + if (result != AllocResult::Success) { + return size_t(0); + } + + *dst_addr = addr; + + try { + KMM_GPU_CHECK(g_memcpy_d_to_h_async( + *dst_addr, + src_addr, + layout.size_in_bytes, + (g_stream_t)stream + )); + return layout.size_in_bytes; + } catch (...) { + m_host_allocator->deallocate_async(stream, addr, layout); + throw; + } + } + ); + + if (result == AllocResult::Success) { + device_state(device_id).stats.bytes_to_host += layout.size_in_bytes; + } + + return result; +} + +AllocResult MemorySystem::allocate_device_and_copy_from_host( + DeviceId device_id, + BufferLayout layout, + g_device_ptr_t* dst_addr, + const void* src_addr, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in, + DeviceEvent& dep_out +) { + AllocResult result; + bool use_hint = + same_context(device_id, stream_hint) && m_events.is_latest_in(stream_hint, deps_in); + + dep_out = schedule_onto_stream( + use_hint, + m_events, + m_streams, + device_id, + StreamKind::HostToDevice, + stream_hint, + deps_in, + [&](const auto& stream) { + void* addr = nullptr; + + result = device_state(device_id).allocator->allocate_async(stream, layout, &addr); + + if (result != AllocResult::Success) { + return size_t(0); + } + + *dst_addr = reinterpret_cast(addr); + + try { + spdlog::trace( + "allocate {} bytes on device {} and copy from host (addr: {}, dst addr: {:#x})", + layout.size_in_bytes, + device_id, + src_addr, + *dst_addr + ); + KMM_GPU_CHECK(g_memcpy_h_to_d_async( + *dst_addr, + src_addr, + layout.size_in_bytes, + (g_stream_t)stream + )); + return layout.size_in_bytes; + } catch (...) { + device_state(device_id).allocator->deallocate(addr, layout); + throw; + } + } + ); + + if (result == AllocResult::Success) { + device_state(device_id).record_allocation(layout.size_in_bytes); + m_host_stats.bytes_to_device[device_id.get()] += layout.size_in_bytes; + } + + return result; +} + +DeviceEvent MemorySystem::fill_device( + DeviceId device_id, + g_device_ptr_t addr, + const FillDescription& description, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in +) { + spdlog::trace("fill device {} memory (addr: {:#x})", device_id, addr); + + return schedule_onto_stream( + same_context(device_id, stream_hint), + m_events, + m_streams, + device_id, + StreamKind::DeviceToDevice, + stream_hint, + deps_in, + [&](g_stream_t stream) -> size_t { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + memops::fill_gpu(stream, reinterpret_cast(addr), description); + return checked_mul(description.num_elements(), description.value.length); +#else + throw std::runtime_error("unsupported operation"); +#endif + } + ); +} + +std::future MemorySystem::fill_host( + void* addr, + const FillDescription& description, + const DeviceEventSet& deps_in +) { + spdlog::trace("fill host memory (addr: {})", addr); + + return std::async(std::launch::async, [this, addr, description, deps_in] { + m_events.synchronize(deps_in); + memops::fill(addr, description); }); } -} // namespace kmm \ No newline at end of file +DeviceEvent MemorySystem::reduce_device( + DeviceId device_id, + g_device_ptr_t src_addr, + g_device_ptr_t dst_addr, + g_device_ptr_t scratch_addr, + const ReductionDescription& description, + const DeviceStreamId& stream_hint, + const DeviceEventSet& deps_in +) { + spdlog::trace( + "reduce device {} memory (src addr: {:#x}, dst addr: {:#x})", + device_id, + src_addr, + dst_addr + ); + + return schedule_onto_stream( + same_context(device_id, stream_hint), + m_events, + m_streams, + device_id, + StreamKind::DeviceToDevice, + stream_hint, + deps_in, + [&](g_stream_t stream) -> memops_extent_type { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + memops::reduce_gpu( + stream, + reinterpret_cast(src_addr), + reinterpret_cast(dst_addr), + reinterpret_cast(scratch_addr), + description + ); + return description.num_outputs(); +#else + throw std::runtime_error("unsupported operation"); +#endif + } + ); +} + +bool MemorySystem::is_copy_supported(MemoryId src, MemoryId dst) const noexcept { + if (src.is_host() || dst.is_host()) { + return true; + } + + auto src_id = src.as_device(); + auto dst_id = dst.as_device(); + + if (src_id == dst_id) { + return true; + } + + return m_peer_access[src_id.get()][dst_id.get()]; +} + +} // namespace kmm diff --git a/src/runtime/reduction_manager.cpp b/src/runtime/reduction_manager.cpp new file mode 100644 index 00000000..c5005c11 --- /dev/null +++ b/src/runtime/reduction_manager.cpp @@ -0,0 +1,523 @@ +#include +#include +#include +#include +#include +#include + +#include "fmt/format.h" + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/flat.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/memops/reduction_gpu.hpp" +#include "kmm/runtime/memory_buffer.hpp" +#include "kmm/runtime/memory_system.hpp" +#include "kmm/runtime/reduction_manager.hpp" + +namespace kmm { + +// A partial buffer handed out by `acquire_partial`: where it lives and the stream it was +// acquired on, so partials can be grouped by location and folded together locally before any +// cross-memory copy. +struct AcquiredPartial { + MemoryBuffer buffer; + MemoryId memory_id; + DeviceStreamId stream_hint; +}; + +/// Backing state for a single reduction: the destination buffer, the op/type/count describing +/// how partials combine, and the partials handed out by `acquire_partial` that still need +/// folding in. Shared (via `ReductionState`) between the caller and the job that does the +/// folding, so it only holds the data everyone cares about -- none of the polling machinery. +class ReductionStateImpl: public reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(ReductionStateImpl) + + public: + ReductionStateImpl(MemoryBuffer home_buffer, ReductionOp op, DataType dtype, size_t count) : + home_buffer(std::move(home_buffer)), + op(op), + dtype(dtype), + count(count) {} + + MemoryBuffer home_buffer; + ReductionOp op; + DataType dtype; + size_t count; + + std::vector partials; + std::unique_ptr job; +}; + +KMM_REFCNT_TRAITS_IMPL(ReductionStateImpl) + +static ReductionDescription describe_reduction( + ReductionOp op, + DataType dtype, + size_t count, + bool accumulate +) { + auto elem_stride = static_cast(data_type_size(dtype)); + auto copy_stride = elem_stride * static_cast(count); + + ReductionDescription description(dtype, op); + description.add_dimension(static_cast(count), elem_stride, elem_stride); + description.reduction_extent = static_cast(1); + description.reduction_stride = copy_stride; + description.accumulate = accumulate; + return description; +} + +// Allocates a fresh `count`-element buffer to accumulate a reduction into (used both by +// `acquire_partial` and by `ReductionJob` when no already-acquired partial can be reused as-is +// as the accumulation target for its group). +static MemoryBuffer allocate_reduction_buffer( + MemoryManager& memory_manager, + DataType dtype, + size_t count, + const std::string& name +) { + size_t elem_size = data_type_size(dtype); + auto layout = BufferLayout {elem_size * count, elem_size}; + auto data = std::make_unique(layout); + return memory_manager.create_buffer(std::move(data), name); +} + +ReductionManager::ReductionManager( + MemoryManager& memory_manager, + refcnt_ptr memory_system +) : + m_memory_manager(memory_manager), + m_memory_system(std::move(memory_system)) {} + +ReductionManager::~ReductionManager() = default; + +ReductionState ReductionManager::initialize_reduction( + MemoryBuffer home_buffer, + ReductionOp op, + DataType dtype, + size_t count +) { + return make_refcnt(std::move(home_buffer), op, dtype, count); +} + +MemoryBuffer ReductionManager::acquire_partial( + ReductionState& reduction, + MemoryId memory_id, + const DeviceStreamId& stream_hint +) { + KMM_ASSERT(!is_submitted(reduction)); + + if (!stream_hint.is_null()) { + for (const auto& partials : reduction->partials) { + if (partials.memory_id == memory_id && partials.stream_hint == stream_hint) { + return partials.buffer; + } + } + } + + auto buffer = allocate_reduction_buffer( + m_memory_manager, + reduction->dtype, + reduction->count, + "reduction-partial:" + reduction->home_buffer->name + ); + + reduction->partials.push_back({buffer, memory_id, stream_hint}); + return buffer; +} + +void ReductionManager::check_compatible( + const ReductionState& reduction, + ReductionOp op, + DataType dtype +) const { + KMM_ASSERT(reduction); + + if (reduction->op != op || reduction->dtype != dtype) { + throw std::runtime_error( + fmt::format( + "reduction contribution does not match the reduction opened by `begin_reduction`: " + "expected op={}, dtype={}, but got op={}, dtype={}", + reduction_op_name(reduction->op), + data_type_name(reduction->dtype), + reduction_op_name(op), + data_type_name(dtype) + ) + ); + } +} + +struct PartialFold { + explicit PartialFold(MemoryId memory_id, MemoryBuffer buffer) : + memory_id(memory_id), + buffer(std::move(buffer)) {} + + Poll poll_ready(MemoryManager& manager, const MemoryTransaction& parent) { + if (!request) { + auto transaction = manager.create_transaction(parent); + request = manager.create_request(buffer, memory_id, AccessKind::ReadOnly, transaction); + } + + return manager.poll_request(DeviceStreamId::null(), request); + } + + void submit( + MemoryManager& manager, + MemorySystem& system, + void* home_addr, + const ReductionDescription& description, + const DeviceEventSet& home_deps + ) { + void* local_addr = manager.access_request(request, request_deps).address; + request_deps.insert(home_deps); + + if (memory_id.is_host()) { + void* peer_addr = local_addr; + host_future = std::async(std::launch::async, [peer_addr, home_addr, description] { + memops::reduce(peer_addr, home_addr, description); + }); + } else { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + // This path does not provide a scratch buffer, so it only supports reductions that + // need none. + KMM_ASSERT(memops::reduce_gpu_scratch_size(description) == 0); + + auto event = system.reduce_device( + memory_id.as_device(), + reinterpret_cast(local_addr), + reinterpret_cast(home_addr), + g_device_ptr_t {}, + description, + DeviceStreamId::null(), + request_deps + ); + + completion_deps.insert(event); +#else + throw std::runtime_error("unsupported operation"); +#endif + } + } + + Poll poll_completed(MemoryManager& manager, DeviceEventSet& deps_out) { + if (host_future.valid()) { + if (host_future.wait_for(std::chrono::seconds(0)) != std::future_status::ready) { + return Poll::Pending; + } + + host_future.get(); + } + + deps_out = release(manager); + return Poll::Ready; + } + + DeviceEventSet release(MemoryManager& manager, DeviceEventSet release_deps = {}) { + if (host_future.valid()) { + host_future.wait(); + } + + release_deps.insert(completion_deps); + manager.release_request(request, release_deps); + request = {}; + done = true; + return release_deps; + } + + MemoryId memory_id; + MemoryBuffer buffer; + MemoryRequest request; + DeviceEventSet request_deps; + bool done = false; + + private: + std::future host_future; + DeviceEventSet completion_deps; +}; + +/// The reduction performed for a single memory. +struct ReductionMemoryJob { + ReductionMemoryJob( + ReductionState reduction, + MemoryId memory_id, + MemoryBuffer home_buffer, + bool home_buffer_initialized, + std::vector buffers = {} + ) : + reduction(reduction), + memory_id(memory_id), + home_buffer(home_buffer), + home_buffer_initialized(home_buffer_initialized) { + partials.reserve(buffers.size()); + + for (auto& buffer : buffers) { + partials.emplace_back(memory_id, std::move(buffer.buffer)); + } + } + + // Partials added this way (e.g. the collapsed result of a per-memory group job) are already + // folded down to a single value per output element. + void add_partial(MemoryBuffer buffer) { + partials.emplace_back(memory_id, std::move(buffer)); + } + + Poll poll(MemoryManager& manager, MemorySystem& system) { + if (active_index) { + try_finalize_active(manager); + } + + if (folded_count == partials.size()) { + if (home_req) { + manager.release_request(home_req, chain_deps); + home_req = {}; + } + + return Poll::Ready; + } + + if (!home_req) { + transaction = manager.create_transaction(); + home_req = + manager.create_request(home_buffer, memory_id, AccessKind::Exclusive, transaction); + } + + DeviceStream stream_hint = {}; + Poll home_ready = manager.poll_request(DeviceStreamId::null(), home_req); + + // Captured once: after the first partial folds into `home_buffer`, further partials + // must wait on `chain_deps` (the previous fold's completion) instead, so `home_deps` is + // deliberately left empty by `try_finalize_active` after that and never repopulated here. + if (home_ready == Poll::Ready && !home_deps_captured) { + home_accessor = manager.access_request(home_req, home_deps); + home_deps_captured = true; + } + + bool work_remaining = active_index.has_value(); + + for (size_t i = 0; i < partials.size(); i++) { + auto& partial = partials[i]; + + if (partial.done || active_index == i) { + continue; + } + + Poll peer_ready = partial.poll_ready(manager, transaction); + + if (peer_ready == Poll::Pending || home_ready == Poll::Pending + || active_index.has_value()) { + work_remaining = true; + continue; + } + + DeviceEventSet wait_deps = chain_deps; + wait_deps.insert(home_deps); + + if (memory_id.is_host() && !wait_deps.is_empty()) { + work_remaining = true; + continue; + } + + auto description = describe_reduction( + reduction->op, + reduction->dtype, + reduction->count, + home_buffer_initialized || folded_count > 0 + ); + + partial.submit(manager, system, home_accessor.address, description, chain_deps); + + active_index = i; + + if (!try_finalize_active(manager)) { + work_remaining = true; + } + } + + return work_remaining ? Poll::Pending : Poll::Ready; + } + + void release(MemoryManager& manager) { + DeviceEventSet deps = chain_deps; + + if (active_index) { + deps = partials[*active_index].release(manager); + active_index.reset(); + } + + if (home_req) { + manager.release_request(home_req, deps); + home_req = {}; + } + + for (auto& partial : partials) { + if (!partial.done && partial.request) { + partial.release(manager); + } + } + + reduction = nullptr; + } + + ReductionState reduction; + MemoryId memory_id; + MemoryBuffer home_buffer; + bool home_buffer_initialized; + size_t folded_count = 0; + MemoryRequest home_req; + MemoryTransaction transaction; + BufferAccessor home_accessor {}; + bool home_deps_captured = false; + DeviceEventSet home_deps; + DeviceEventSet chain_deps; + std::vector partials; + std::optional active_index; + + private: + bool try_finalize_active(MemoryManager& manager) { + DeviceEventSet release_deps; + + if (partials[*active_index].poll_completed(manager, release_deps) == Poll::Pending) { + return false; + } + + home_deps.clear(); + chain_deps = release_deps; + folded_count++; + active_index.reset(); + return true; + } +}; + +struct ReductionJob { + ReductionJob( + ReductionState reduction, + MemoryId memory_id, + MemoryBuffer home_buffer, + std::vector buffers, + MemoryManager& memory_manager, + refcnt_ptr memory_system + ) { + std::stable_sort(buffers.begin(), buffers.end(), [](const auto& a, const auto& b) { + return a.memory_id < b.memory_id; + }); + + size_t index = 0; + while (index < buffers.size()) { + auto group_memory_id = buffers[index].memory_id; + + std::vector partials; + MemoryBuffer group_home_buffer = buffers[index].buffer; + bool group_home_initialized = true; + index++; + + while (index < buffers.size() && buffers[index].memory_id == group_memory_id) { + partials.push_back(buffers[index]); + index++; + } + + group_jobs.push_back( + std::make_unique( + reduction, + group_memory_id, + std::move(group_home_buffer), + group_home_initialized, + std::move(partials) + ) + ); + } + + final_reduction = std::make_unique( + reduction, + memory_id, + std::move(home_buffer), + /* home_buffer_initialized = */ false + ); + } + + Poll poll(MemoryManager& manager, MemorySystem& system) { + Poll result = Poll::Ready; + + for (auto& job : group_jobs) { + if (!job) { + // Already finished and handed off to `final_reduction` on an earlier tick. + continue; + } + + if (job->poll(manager, system) == Poll::Pending) { + result = Poll::Pending; + continue; + } + + final_reduction->add_partial(std::move(job->home_buffer)); + job = nullptr; + } + + if (final_reduction->poll(manager, system) == Poll::Pending) { + result = Poll::Pending; + } + + return result; + } + + void release(MemoryManager& manager) { + for (auto& job : group_jobs) { + if (job) { + job->release(manager); + job = nullptr; + } + } + + final_reduction->release(manager); + final_reduction = nullptr; + + for (auto& buffer : owned_buffers) { + manager.release_buffer(buffer); + } + + owned_buffers.clear(); + } + + private: + std::vector> group_jobs; + std::unique_ptr final_reduction; + std::vector owned_buffers; +}; + +void ReductionManager::submit_reduction(ReductionState& reduction, MemoryId memory_id) { + KMM_ASSERT(!is_submitted(reduction)); + reduction->job = std::make_unique( + reduction, + memory_id, + reduction->home_buffer, + reduction->partials, + m_memory_manager, + m_memory_system + ); +} + +bool ReductionManager::is_submitted(const ReductionState& reduction) { + return bool(reduction) && reduction->job != nullptr; +} + +Poll ReductionManager::poll_reduction(ReductionState& reduction) { + KMM_ASSERT(reduction->job); + return reduction->job->poll(m_memory_manager, *m_memory_system); +} + +void ReductionManager::release_reduction(ReductionState& reduction) { + // Release any request the job is still holding. + if (reduction->job) { + reduction->job->release(m_memory_manager); + reduction->job = nullptr; + } + + // Every partial ever handed out by `acquire_partial` is still tracked here, regardless of + // whether/how the job above grouped and folded it, so release them all directly. + for (auto& partial : reduction->partials) { + m_memory_manager.release_buffer(partial.buffer); + } +} + +} // namespace kmm diff --git a/src/runtime/runtime.cpp b/src/runtime/runtime.cpp index 94a3f2c5..2520fa4f 100644 --- a/src/runtime/runtime.cpp +++ b/src/runtime/runtime.cpp @@ -1,298 +1,933 @@ +#include +#include #include - -#include "fmt/std.h" -#include "spdlog/spdlog.h" - -#include "kmm/runtime/allocators/block.hpp" -#include "kmm/runtime/allocators/caching.hpp" -#include "kmm/runtime/allocators/device.hpp" -#include "kmm/runtime/allocators/system.hpp" +#include + +#include "kmm/core/macros.hpp" +#include "kmm/core/panic.hpp" +#include "kmm/runtime/data_interfaces/external.hpp" +#include "kmm/runtime/data_interfaces/flat.hpp" +#include "kmm/runtime/data_interfaces/managed.hpp" +#include "kmm/runtime/data_interfaces/pinned.hpp" +#include "kmm/runtime/device_data_streams.hpp" +#include "kmm/runtime/memops/reduction_gpu.hpp" +#include "kmm/runtime/memory_buffer.hpp" +#include "kmm/runtime/reduction_manager.hpp" +#include "kmm/runtime/resource.hpp" #include "kmm/runtime/runtime.hpp" +#include "kmm/utils/scope_exit.hpp" namespace kmm { -static SystemInfo make_system_info( - const std::vector& contexts, - const RuntimeConfig& config -) { - spdlog::info("detected {} GPU device(s):", contexts.size()); - std::vector device_infos; +struct BufferEntry { + KMM_NOT_COPYABLE(BufferEntry) - for (size_t i = 0; i < contexts.size(); i++) { - auto info = DeviceInfo(DeviceId(i), contexts[i], config.device_concurrent_streams); - auto memory_gb = static_cast(info.total_memory_size()) / 1e9; + public: + MemoryBuffer buffer; + std::exception_ptr poison = nullptr; + ReductionState reduction; + DeviceEventSet last_accesses; - spdlog::info(" - {} ({:.2} GB)", info.name(), memory_gb); - device_infos.push_back(info); + BufferEntry(MemoryBuffer&& buffer) : buffer(std::move(buffer)) {} + + void record_access(const DeviceEventSet& deps) { + last_accesses.insert(deps); + } +}; + +/// Owns everything a `Runtime` hands out: the machine topology, the stream registry, the +/// physical-memory backend, and the `MemoryManager` built on top of it. +class RuntimeImpl: public reference_count { + KMM_NOT_COPYABLE_OR_MOVABLE(RuntimeImpl) + + public: + RuntimeImpl(const RuntimeConfig& config) : + default_buffer_kind(config.default_buffer_kind), + system_info {}, + memory_system {std::make_unique(system_info, event_registry, config)}, + memory_manager(memory_system), + reduction_manager(memory_manager, memory_system) {} + + // Buffers are normally released one at a time via `Runtime::release_buffer`, but any that + // are still outstanding when the runtime itself is torn down need to be released here too: + // otherwise `buffers` would just drop its `MemoryBuffer` refs, and their host/device + // allocations would never go through `MemoryManager::release_buffer`. + ~RuntimeImpl() { + for (auto& [id, entry] : buffers) { + if (entry.reduction) { + reduction_manager.release_reduction(entry.reduction); + } + + memory_manager.release_buffer(std::move(entry.buffer)); + } } - return device_infos; -} + BufferEntry& find_entry(BufferId id) { + auto it = buffers.find(id); + + if (it == buffers.end()) { + throw std::runtime_error(fmt::format("could not find buffer {}", id.get())); + } + + if (it->second.poison) { + std::rethrow_exception(it->second.poison); + } + + return it->second; + } + + MemoryBuffer find_buffer(BufferId id) { + return find_entry(id).buffer; + } + + void poison_buffer(BufferId id, std::exception_ptr reason) noexcept { + auto it = buffers.find(id); + + if (it != buffers.end() && !it->second.poison) { + it->second.poison = std::move(reason); + } + } + + std::chrono::system_clock::time_point poll_once() { + static constexpr auto poll_interval = std::chrono::milliseconds(5); + auto now = std::chrono::system_clock::now(); + + if (now >= next_poll_deadline) { + event_registry.make_progress(); + memory_system->make_progress(); + memory_manager.make_progress(); + + next_poll_deadline = now + poll_interval; + } + + return next_poll_deadline; + } + + template + void poll_until_completion(std::unique_lock& guard, F callback) { + if (callback()) { + return; + } + + auto deadline = poll_once(); + + while (!callback()) { + guard.unlock(); + std::this_thread::sleep_until(deadline); + guard.lock(); + + deadline = poll_once(); + } + } + + const BufferKind default_buffer_kind; + const SystemInfo system_info; -Runtime::Runtime( - std::vector contexts, - std::shared_ptr stream_manager, - std::shared_ptr memory_system, - const RuntimeConfig& config -) : - m_memory_system(memory_system), - m_memory_manager(std::make_shared(memory_system)), - m_buffer_registry(std::make_shared(m_memory_manager)), - m_stream_manager(stream_manager), - m_devices(std::make_shared( - contexts, - config.device_concurrent_streams, - m_stream_manager - )), - m_info(make_system_info(contexts, config)), - m_scheduler(m_devices, stream_manager, m_buffer_registry, config.debug_mode) {} + std::mutex mutex; + DeviceEventRegistry event_registry; + refcnt_ptr memory_system; + MemoryManager memory_manager; + ReductionManager reduction_manager; + ankerl::unordered_dense::map buffers; + uint64_t next_buffer_id_counter = 0; + std::chrono::system_clock::time_point next_poll_deadline; +}; -Runtime::~Runtime() { - shutdown(); +KMM_REFCNT_TRAITS_IMPL(RuntimeImpl) + +Runtime::Runtime(refcnt_ptr impl) : m_impl(std::move(impl)) { + KMM_ASSERT(m_impl); } -BufferId Runtime::create_buffer(BufferLayout layout) { - return this->schedule([&](TaskGraph& g) { // - return g.create_buffer(layout); - }); +MemorySystem& Runtime::memory_system() noexcept { + return *m_impl->memory_system; } -void Runtime::delete_buffer(BufferId id, EventList deps) { - this->schedule([&](TaskGraph& g) { // - g.delete_buffer(id, std::move(deps)); - }); +DeviceEventRegistry& Runtime::event_registry() noexcept { + return m_impl->event_registry; } -void Runtime::check_buffer(BufferId id) { - std::unique_lock guard {m_mutex}; - m_buffer_registry->get(id); +const SystemInfo& Runtime::system_info() const noexcept { + return m_impl->system_info; } -bool Runtime::query_event(EventId event_id, std::chrono::system_clock::time_point deadline) { - std::unique_lock guard {m_mutex}; - make_progress_impl(); +std::chrono::system_clock::time_point Runtime::poll_once() { + std::lock_guard guard {m_impl->mutex}; + return m_impl->poll_once(); +} - while (!m_scheduler.is_completed(event_id)) { - KMM_ASSERT(!m_scheduler.is_idle()); - auto next_update = m_next_updated_planned; +bool Runtime::poll_until_completion( + function_ref callback, + const std::chrono::system_clock::time_point deadline +) { + auto next_poll_at = poll_once(); - if (next_update > deadline) { + while (!callback()) { + // Next poll is after the deadline. Return false now instead waiting. + if (next_poll_at > deadline) { return false; } - guard.unlock(); - std::this_thread::sleep_until(next_update); - guard.lock(); - - make_progress_impl(); + std::this_thread::sleep_until(next_poll_at); + next_poll_at = poll_once(); } return true; } -bool Runtime::is_idle() { - std::lock_guard guard {m_mutex}; - return is_idle_impl(); +void Runtime::synchronize(const DeviceEvent& e) { + poll_until_completion([&] { return m_impl->event_registry.is_ready(e); }); } -void Runtime::trim_memory() { - std::lock_guard guard {m_mutex}; - m_memory_system->trim_host(); - m_memory_system->trim_device(); +void Runtime::synchronize(const DeviceEventSet& events) { + for (const auto& e : events) { + if (!m_impl->event_registry.is_ready(e)) { + synchronize(e); + } + } } -void Runtime::make_progress() { - std::lock_guard guard {m_mutex}; - make_progress_impl(); +void Runtime::synchronize() { + poll_until_completion([&] { return m_impl->event_registry.is_all_ready(); }); } -void Runtime::shutdown() { - std::unique_lock guard {m_mutex}; - if (m_has_shutdown) { - return; +BufferId Runtime::register_buffer( + std::unique_ptr data, + std::string name, + std::optional home, + bool evictable +) { + if (data == nullptr) { + throw std::runtime_error("register_buffer: data interface must not be null"); } - m_has_shutdown = true; + std::lock_guard guard(m_impl->mutex); + + auto id = BufferId(m_impl->next_buffer_id_counter++); + + if (name.empty()) { + name = std::to_string(id.get()); + } - while (!is_idle_impl()) { - make_progress_impl(); + auto buffer = + m_impl->memory_manager.create_buffer(std::move(data), std::move(name), evictable, home); + m_impl->buffers.emplace(id, std::move(buffer)); + return id; +} - guard.unlock(); - std::this_thread::sleep_for(std::chrono::milliseconds {10}); - guard.lock(); +BufferId Runtime::create_buffer( + BufferLayout layout, + std::string name, + FillValue fill_value, + std::optional home, + std::optional kind +) { + std::unique_ptr data; + switch (kind.value_or(m_impl->default_buffer_kind)) { + case BufferKind::Discrete: + data = std::make_unique(layout, std::move(fill_value)); + break; + case BufferKind::Managed: + data = std::make_unique(layout, std::move(fill_value)); + break; + case BufferKind::HostPinned: + if (fill_value.length != 0) { + throw std::runtime_error( + "create_buffer: BufferKind::HostPinned does not support a fill value" + ); + } + data = std::make_unique(layout); + break; } - m_stream_manager->wait_until_idle(); + return register_buffer(std::move(data), std::move(name), home, /* evictable = */ true); } -EventId Runtime::commit_impl(TaskGraph& g) { - std::vector nodes_out; - std::vector> buffers_out; +BufferId Runtime::adopt_buffer( + BufferLayout layout, + std::string name, + void* external_ptr, + MemoryId memory_id +) { + auto data = + std::make_unique(external_ptr, layout.size_in_bytes, memory_id); + + // Not evictable: KMM does not own the allocation and cannot recreate it after eviction. + return register_buffer( + std::move(data), + std::move(name), + memory_id, + /* evictable = */ false + ); +} + +void Runtime::prefetch_buffer(BufferId id, MemoryId memory_id, bool invalidate_others) { + std::lock_guard guard(m_impl->mutex); + AccessKind a = invalidate_others ? AccessKind::Exclusive : AccessKind::ReadOnly; + m_impl->memory_manager.prefetch_buffer(m_impl->find_buffer(id), memory_id, a); +} + +void Runtime::poison_buffer(BufferId id, std::exception_ptr reason) noexcept { + std::lock_guard guard(m_impl->mutex); + m_impl->poison_buffer(id, std::move(reason)); +} + +bool Runtime::is_valid(BufferId id, MemoryId memory_id) { + std::lock_guard guard(m_impl->mutex); + return m_impl->find_buffer(id)->is_valid(memory_id); +} - auto barrier_id = m_graph_state.commit(g, nodes_out, buffers_out); +std::optional Runtime::find_valid_memory(BufferId id) const { + std::lock_guard guard(m_impl->mutex); + const auto& buf = m_impl->find_buffer(id); - // Flush all staged buffers to the registry - for (auto&& [id, layout] : buffers_out) { - m_buffer_registry->add(id, layout); + if (buf->is_valid(MemoryId::host())) { + return MemoryId::host(); } - // Flush all events from the DAG builder to the scheduler - for (auto&& e : nodes_out) { - m_scheduler.submit( - e.id, // - build_task_for_command(std::move(e.command)), - std::move(e.dependencies) - ); + for (size_t i = 0; i < MAX_DEVICES; i++) { + if (buf->is_valid(MemoryId::device(DeviceId(i)))) { + return MemoryId::device(DeviceId(i)); + } } - // Plan an update to happen now since we have added new tasks to the scheduler. - m_next_updated_planned = std::chrono::system_clock::time_point::min(); + return std::nullopt; +} + +std::optional Runtime::buffer_home(BufferId id) const { + std::lock_guard guard(m_impl->mutex); + return m_impl->find_buffer(id)->home_memory_id; +} + +bool Runtime::is_allocated(BufferId id, MemoryId memory_id) { + std::lock_guard guard(m_impl->mutex); + return m_impl->find_buffer(id)->is_allocated(memory_id); +} - return barrier_id; +void Runtime::try_evict_buffer(BufferId id, MemoryId memory_id) { + std::lock_guard guard(m_impl->mutex); + m_impl->memory_manager.try_evict_buffer(m_impl->find_buffer(id), memory_id); } -void Runtime::make_progress_impl() { - static constexpr auto TIMEOUT = std::chrono::microseconds {100}; - auto now = std::chrono::system_clock::now(); +void Runtime::invalidate_buffer(BufferId id) { + std::lock_guard guard(m_impl->mutex); + m_impl->memory_manager.invalidate_buffer(m_impl->find_buffer(id)); +} - if (m_next_updated_planned > now) { - return; +void Runtime::trim(MemoryId memory_id, size_t bytes_to_keep, bool evict) { + std::lock_guard guard(m_impl->mutex); + + if (memory_id.is_host()) { + // No eviction path for host memory; `evict` only affects device memory. + m_impl->memory_system->trim_host(bytes_to_keep); + } else { + m_impl->memory_manager.trim_device(memory_id.as_device(), bytes_to_keep, evict); + } +} + +void Runtime::begin_reduction(BufferId id, DataType dtype, ReductionOp op) { + std::lock_guard guard(m_impl->mutex); + auto& entry = m_impl->find_entry(id); + + if (entry.reduction) { + throw std::runtime_error(fmt::format("buffer {} is already in reduction mode", id.get())); } - m_next_updated_planned = now + TIMEOUT; - m_stream_manager->make_progress(); - m_memory_system->make_progress(); - m_scheduler.make_progress(); + // The number of output elements the partials and the fold operate on is derived from the + // buffer size: `begin_reduction` fixes the element type, so the buffer holds exactly this + // many values. + size_t elem_size = data_type_size(dtype); + KMM_ASSERT(elem_size > 0 && entry.buffer->size_in_bytes % elem_size == 0); + size_t count = entry.buffer->size_in_bytes / elem_size; + + entry.reduction = + m_impl->reduction_manager.initialize_reduction(entry.buffer, op, dtype, count); } -bool Runtime::is_idle_impl() { - return m_stream_manager->is_idle() && m_scheduler.is_idle() - && m_memory_manager->is_idle(*m_stream_manager); +void Runtime::finalize_reduction(BufferId id, MemoryId memory_id) { + std::unique_lock guard(m_impl->mutex); + auto& entry = m_impl->find_entry(id); + auto reduction = entry.reduction; + + if (!reduction) { + throw std::runtime_error(fmt::format("buffer {} is not in reduction mode", id.get())); + } + + m_impl->reduction_manager.submit_reduction(reduction, memory_id); + + try { + m_impl->poll_until_completion(guard, [&] { + return m_impl->reduction_manager.poll_reduction(reduction) == Poll::Ready; + }); + } catch (...) { + m_impl->reduction_manager.release_reduction(reduction); + entry.reduction = nullptr; + throw; + } + + m_impl->reduction_manager.release_reduction(reduction); + entry.reduction = nullptr; } -static size_t compute_device_memory_limit( - const RuntimeConfig& config, - const GPUContextHandle& context -) { - // ignore `device_memory_reserved` if it is zero - if (config.device_memory_keep_free == 0) { - return config.device_memory_limit; +void Runtime::rollback_reduction(BufferId id) { + std::lock_guard guard(m_impl->mutex); + auto& entry = m_impl->find_entry(id); + + if (!entry.reduction) { + return; } - GPUContextGuard guard {context}; + m_impl->reduction_manager.release_reduction(entry.reduction); + entry.reduction = nullptr; +} - size_t memory_capacity, memory_available; - KMM_GPU_CHECK(g_mem_get_info(&memory_available, &memory_capacity)); +void Runtime::release_buffer(BufferId id) { + std::lock_guard guard(m_impl->mutex); - // Insufficient memory capacity - if (memory_capacity < config.device_memory_keep_free) { - spdlog::warn( - "cannot keep {} bytes available on GPU, memory capacity is only {} bytes", - config.device_memory_keep_free, - memory_capacity - ); + auto it = m_impl->buffers.find(id); + auto& entry = it->second; - return 0; + if (entry.reduction) { + // If the reduction was submitted, we need to release it. + m_impl->reduction_manager.release_reduction(entry.reduction); } - return std::min( - memory_capacity - config.device_memory_keep_free, // - config.device_memory_limit - ); + m_impl->memory_manager.release_buffer(std::move(entry.buffer)); + m_impl->buffers.erase(it); } -std::unique_ptr create_device_allocator( - const RuntimeConfig& config, - const GPUContextHandle& context, - std::shared_ptr stream_manager +ResourceGrant Runtime::submit( + ResourceRequest requests, + std::optional stream, + MemoryTransaction parent ) { - std::unique_ptr alloc; - size_t memory_limit = compute_device_memory_limit(config, context); + std::unique_lock guard(m_impl->mutex); - switch (config.device_memory_kind) { - case DeviceMemoryKind::NoPool: - return std::make_unique(context, stream_manager, memory_limit); - ; + auto stream_id = DeviceStreamId::null(); - case DeviceMemoryKind::CachingPool: - alloc = std::make_unique(context, stream_manager, memory_limit); + if (stream.has_value()) { + stream_id = m_impl->event_registry.lookup_or_register_stream(*stream); + } - return std::make_unique(std::move(alloc)); + auto transaction = m_impl->memory_manager.create_transaction(parent); - case DeviceMemoryKind::DefaultPool: - return std::make_unique( - context, - stream_manager, - DevicePoolKind::Default, - memory_limit - ); + std::vector entries; + entries.reserve(requests.m_requests.size()); - case DeviceMemoryKind::PrivatePool: - return std::make_unique( - context, - stream_manager, - DevicePoolKind::Create, - memory_limit - ); + for (const auto& req : requests.m_requests) { + entries.push_back(ResourceGrant::Entry {req.buffer_id, nullptr, {}}); + } - default: - KMM_PANIC("invalid memory kind"); + for (size_t i = 0; i < entries.size(); i++) { + auto& it = entries[i]; + const auto& req = requests.m_requests[i]; + + try { + auto& entry = m_impl->find_entry(it.buffer_id); + + if (req.mode == AccessMode::Reduce) { + if (!entry.reduction) { + throw std::runtime_error( + fmt::format( + "buffer {} is not in reduction mode, call `begin_reduction` before " + "reducing into it", + it.buffer_id.get() + ) + ); + } + + auto partial = m_impl->reduction_manager.acquire_partial( // + entry.reduction, + req.memory_id, + stream_id + ); + + it.request = m_impl->memory_manager.create_request( // + partial, + req.memory_id, + AccessKind::Exclusive, + transaction + ); + } else { + if (entry.reduction) { + throw std::runtime_error( + fmt::format( + "buffer {} is still in reduction mode, call `finalize_reduction` " + "before reading or writing it", + it.buffer_id.get() + ) + ); + } + + auto access = + req.mode == AccessMode::Read ? AccessKind::ReadOnly : AccessKind::Exclusive; + + it.request = m_impl->memory_manager.create_request( // + entry.buffer, + req.memory_id, + access, + transaction + ); + } + } catch (...) { + for (auto& rollback : entries) { + if (rollback.request) { + m_impl->memory_manager.release_request(std::move(rollback.request)); + } + } + + throw; + } } + + DeviceEventSet deps; + + try { + m_impl->poll_until_completion(guard, [&] { + bool is_ready = true; + + for (auto& entry : entries) { + if (m_impl->memory_manager.poll_request(stream_id, entry.request) + == Poll::Pending) { + is_ready = false; + } + } + + return is_ready; + }); + + for (auto& entry : entries) { + entry.accessor = m_impl->memory_manager.access_request(entry.request, deps); + } + + if (stream_id.is_null()) { + for (const auto& dep : deps) { + m_impl->poll_until_completion(guard, [&] { + return m_impl->event_registry.is_ready(dep); + }); + } + } else { + m_impl->event_registry.wait_on_event(stream_id, deps); + } + } catch (...) { + for (auto& entry : entries) { + if (entry.request) { + m_impl->memory_manager.release_request(std::move(entry.request), {}); + } + } + + throw; + } + + return ResourceGrant(std::move(entries), std::move(deps), std::move(transaction)); } -std::shared_ptr make_worker(const RuntimeConfig& config) { - std::unique_ptr host_mem; - std::vector> device_mems; +void Runtime::release(ResourceGrant& grant, DeviceEventSet deps) { + std::lock_guard guard(m_impl->mutex); - auto stream_manager = std::make_shared(); - auto contexts = std::vector(); - auto devices = get_gpu_devices(); + deps.insert(grant.m_deps); - if (devices.empty()) { - host_mem = std::make_unique(stream_manager, config.host_memory_limit); - } else if (devices.size() > MAX_DEVICES) { - throw std::runtime_error(fmt::format("cannot support more than {} GPU(s)", MAX_DEVICES)); - } else { - for (const auto& device : devices) { - auto context = GPUContextHandle::retain_primary_context_for_device(device); - device_mems.push_back(create_device_allocator(config, context, stream_manager)); - contexts.push_back(std::move(context)); + for (auto& entry : grant.m_entries) { + m_impl->find_entry(entry.buffer_id).record_access(deps); + m_impl->memory_manager.release_request(std::move(entry.request), deps); + } + + grant.m_entries.clear(); + grant.m_deps.clear(); + grant.m_transaction = {}; +} + +void Runtime::poison(const ResourceGrant& grant, std::exception_ptr reason) noexcept { + for (const auto& entry : grant.m_entries) { + if (entry.accessor.is_writable) { + poison_buffer(entry.buffer_id, reason); } + } +} - host_mem = std::make_unique( - contexts.at(0), - stream_manager, - config.host_memory_limit +static DeviceEvent do_copy( + std::unique_lock& guard, + RuntimeImpl* impl, + BufferAccessor dst_access, + BufferAccessor src_access, + const CopyDescription& description, + const DeviceEventSet& deps, + const DeviceStreamId& stream_hint = {} +) { + KMM_ASSERT(range(dst_access.size_in_bytes).contains(description.dst_range())); + KMM_ASSERT(range(src_access.size_in_bytes).contains(description.src_range())); + + auto dst_memory_id = dst_access.memory_id; + auto src_memory_id = src_access.memory_id; + + DeviceEvent event; + + if (dst_memory_id.is_device() && src_memory_id.is_device()) { + event = impl->memory_system->copy_device_to_device( + src_memory_id.as_device(), + dst_memory_id.as_device(), + g_device_ptr_t(static_cast(src_access.address) + description.src_offset), + g_device_ptr_t(static_cast(dst_access.address) + description.dst_offset), + description.element_size, + stream_hint, + deps + ); + } else if (dst_memory_id.is_device() && src_memory_id.is_host()) { + event = impl->memory_system->copy_host_to_device( + dst_memory_id.as_device(), + reinterpret_cast(src_access.address) + description.src_offset, + g_device_ptr_t(static_cast(dst_access.address) + description.dst_offset), + description.element_size, + stream_hint, + deps ); + } else if (dst_memory_id.is_host() && src_memory_id.is_device()) { + event = impl->memory_system->copy_device_to_host( + src_memory_id.as_device(), + g_device_ptr_t(static_cast(src_access.address) + description.src_offset), + reinterpret_cast(dst_access.address) + description.dst_offset, + description.element_size, + stream_hint, + deps + ); + } else { + // TOOD: maybe use a thread pool? + auto fut = + std::async([=] { memops::copy(src_access.address, dst_access.address, description); }); + + impl->poll_until_completion(guard, [&] { + if (fut.wait_for(std::chrono::seconds(0)) == std::future_status::timeout) { + return false; + } + + fut.get(); + return true; + }); + } + + return event; +} + +DeviceEvent Runtime::submit_copy( + BufferId dst_id, + BufferId src_id, + CopyDescription description, + MemoryId memory_id, + std::optional stream, + MemoryTransaction parent +) { + std::unique_lock guard(m_impl->mutex); + description = description.simplify(); + + auto stream_hint = DeviceStreamId::null(); + + if (stream.has_value()) { + stream_hint = m_impl->event_registry.lookup_or_register_stream(*stream); + } + + auto& dst_entry = m_impl->find_entry(dst_id); + auto& src_entry = m_impl->find_entry(src_id); + + bool same_buffer = dst_id == src_id; + + MemoryId dst_memory_id = memory_id; + MemoryId src_memory_id = + same_buffer ? dst_memory_id : src_entry.buffer->find_preferred_location(memory_id); + + auto transaction = m_impl->memory_manager.create_transaction(parent); + + MemoryRequest dst_req = m_impl->memory_manager.create_request( // + dst_entry.buffer, + dst_memory_id, + AccessKind::Exclusive, + transaction + ); + + MemoryRequest src_req = dst_req; + + if (!same_buffer) { + try { + src_req = m_impl->memory_manager.create_request( // + src_entry.buffer, + src_memory_id, + AccessKind::ReadOnly, + transaction + ); + } catch (...) { + // we must release the other request as well. + m_impl->memory_manager.release_request(std::move(dst_req)); + throw; + } + } - if (config.host_memory_kind == HostMemoryKind::CachingPool) { - host_mem = std::make_unique(std::move(host_mem)); + DeviceEventSet deps; + DeviceEvent event; + + try { + m_impl->poll_until_completion(guard, [&] { + auto dst_status = m_impl->memory_manager.poll_request(stream_hint, dst_req); + auto src_status = m_impl->memory_manager.poll_request(stream_hint, src_req); + return src_status == Poll::Ready && dst_status == Poll::Ready; + }); + + auto dst_accessor = m_impl->memory_manager.access_request(dst_req, deps); + auto src_accessor = m_impl->memory_manager.access_request(src_req, deps); + + event = do_copy( + guard, + m_impl.get(), + dst_accessor, + src_accessor, + description, + deps, + stream_hint + ); + } catch (...) { + if (!same_buffer) { + m_impl->memory_manager.release_request(src_req); + m_impl->memory_manager.release_request(dst_req); + } else { + m_impl->memory_manager.release_request(dst_req); } + + throw; + } + + if (!same_buffer) { + m_impl->memory_manager.release_request(src_req, event); + m_impl->memory_manager.release_request(dst_req, event); + } else { + m_impl->memory_manager.release_request(dst_req, event); } - if (config.host_memory_block_size > 0) { - host_mem = std::make_unique( // - std::move(host_mem), - config.host_memory_block_size + return event; +} + +static DeviceEvent do_reduce( + std::unique_lock& guard, + RuntimeImpl* impl, + BufferAccessor dst_access, + BufferAccessor src_access, + BufferAccessor scratch_access, + const ReductionDescription& description, + const DeviceEventSet& deps, + const DeviceStreamId& stream_hint = {} +) { + KMM_ASSERT(range(dst_access.size_in_bytes).contains(description.dst_range())); + KMM_ASSERT(range(src_access.size_in_bytes).contains(description.src_range())); + + auto memory_id = dst_access.memory_id; + DeviceEvent event; + + if (memory_id.is_device()) { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + KMM_ASSERT(scratch_access.size_in_bytes >= memops::reduce_gpu_scratch_size(description)); + + event = impl->memory_system->reduce_device( + memory_id.as_device(), + g_device_ptr_t(src_access.address), + g_device_ptr_t(dst_access.address), + g_device_ptr_t(scratch_access.address), + description, + stream_hint, + deps ); +#else + throw std::runtime_error("unsupported operation"); +#endif + } else { + // TODO: maybe use a thread pool? + auto fut = std::async([=] { + memops::reduce(src_access.address, dst_access.address, description); + }); + + impl->poll_until_completion(guard, [&] { + if (fut.wait_for(std::chrono::seconds(0)) == std::future_status::timeout) { + return false; + } + + fut.get(); + return true; + }); } - if (config.device_memory_block_size > 0) { - for (size_t i = 0; i < devices.size(); i++) { - device_mems[i] = std::make_unique( - std::move(device_mems[i]), - config.device_memory_block_size + return event; +} + +DeviceEvent Runtime::submit_reduction( + BufferId dst_id, + BufferId src_id, + ReductionDescription description, + MemoryId memory_id, + std::optional stream, + MemoryTransaction parent +) { + std::unique_lock guard(m_impl->mutex); + + auto stream_hint = DeviceStreamId::null(); + + if (stream.has_value()) { + stream_hint = m_impl->event_registry.lookup_or_register_stream(*stream); + } + + auto& dst_entry = m_impl->find_entry(dst_id); + auto& src_entry = m_impl->find_entry(src_id); + + if (src_entry.reduction) { + throw std::runtime_error( + fmt::format( + "buffer {} is still in reduction mode and cannot be a reduction source, call " + "`finalize_reduction` before reading it", + src_id.get() + ) + ); + } + + // If the destination is mid-reduction its home buffer must not be written directly (the + // outstanding partials have not been folded in yet). Route this contribution into a fresh + // partial of the ongoing reduction instead; `finalize_reduction` folds it in with the rest. + MemoryBuffer dst_buffer = dst_entry.buffer; + + if (dst_entry.reduction) { + if (m_impl->reduction_manager.is_submitted(dst_entry.reduction)) { + throw std::runtime_error( + fmt::format( + "buffer {} reduction has already been finalized, cannot reduce into it", + dst_id.get() + ) ); } + + m_impl->reduction_manager.check_compatible( + dst_entry.reduction, + description.operation, + description.dtype + ); + + dst_buffer = + m_impl->reduction_manager.acquire_partial(dst_entry.reduction, memory_id, stream_hint); + + // The partial is freshly allocated: there is nothing in it to combine with. + description.accumulate = false; } - auto memory_system = std::make_shared( - stream_manager, - contexts, - std::move(host_mem), - std::move(device_mems) + MemoryBuffer scratch_buffer; + MemoryRequest scratch_req; + MemoryRequest dst_req; + MemoryRequest src_req; + + // if an exception occurs, then we need to clean up the memory requests and scratch buffer. + auto cleanup = scope_exit([&] { + if (dst_req) { + m_impl->memory_manager.release_request(dst_req); + } + + if (src_req) { + m_impl->memory_manager.release_request(src_req); + } + + if (scratch_req) { + m_impl->memory_manager.release_request(scratch_req); + } + + if (scratch_buffer) { + m_impl->memory_manager.release_buffer(scratch_buffer); + } + }); + +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + size_t scratch_size = memops::reduce_gpu_scratch_size(description); +#else + size_t scratch_size = 0; +#endif + bool has_scratch = memory_id.is_device() && scratch_size > 0; + + auto transaction = m_impl->memory_manager.create_transaction(parent); + + dst_req = m_impl->memory_manager.create_request( // + dst_buffer, + memory_id, + AccessKind::Exclusive, + transaction ); - return std::make_shared(contexts, stream_manager, memory_system, config); + src_req = m_impl->memory_manager.create_request( // + src_entry.buffer, + memory_id, + AccessKind::ReadOnly, + transaction + ); + + if (has_scratch) { + auto layout = BufferLayout::for_type(scratch_size); + auto iface = std::make_unique(layout); + + scratch_buffer = m_impl->memory_manager.create_buffer( + std::move(iface), + "reduction-scratch-buffer", + false, + memory_id + ); + + scratch_req = m_impl->memory_manager.create_request( // + scratch_buffer, + memory_id, + AccessKind::Exclusive, + transaction + ); + } + + DeviceEventSet deps; + DeviceEvent event; + + m_impl->poll_until_completion(guard, [&] { + auto dst_status = m_impl->memory_manager.poll_request(stream_hint, dst_req); + auto src_status = m_impl->memory_manager.poll_request(stream_hint, src_req); + auto scratch_status = scratch_req + ? m_impl->memory_manager.poll_request(stream_hint, scratch_req) + : Poll::Ready; + return src_status == Poll::Ready && dst_status == Poll::Ready + && scratch_status == Poll::Ready; + }); + + auto dst_accessor = m_impl->memory_manager.access_request(dst_req, deps); + auto src_accessor = m_impl->memory_manager.access_request(src_req, deps); + + BufferAccessor scratch_accessor {}; + if (scratch_req) { + scratch_accessor = m_impl->memory_manager.access_request(scratch_req, deps); + } + + event = do_reduce( + guard, + m_impl.get(), + dst_accessor, + src_accessor, + scratch_accessor, + description, + deps, + stream_hint + ); + + m_impl->memory_manager.release_request(src_req, event); + src_req = nullptr; + + m_impl->memory_manager.release_request(dst_req, event); + dst_req = nullptr; + + if (has_scratch) { + m_impl->memory_manager.release_request(scratch_req, event); + scratch_req = nullptr; + + m_impl->memory_manager.release_buffer(scratch_buffer); + scratch_buffer = nullptr; + } + + return event; } + +Runtime make_runtime(const RuntimeConfig& config) { + return Runtime(std::make_unique(config)); +} + } // namespace kmm diff --git a/src/core/config.cpp b/src/runtime/runtime_config.cpp similarity index 97% rename from src/core/config.cpp rename to src/runtime/runtime_config.cpp index 8a66fa93..64e3da03 100644 --- a/src/core/config.cpp +++ b/src/runtime/runtime_config.cpp @@ -7,8 +7,8 @@ #include "fmt/format.h" #include "spdlog/spdlog.h" -#include "kmm/core/config.hpp" -#include "kmm/utils/checked_math.hpp" +#include "kmm/core/checked_math.hpp" +#include "kmm/runtime/runtime_config.hpp" namespace kmm { diff --git a/src/runtime/scheduler.cpp b/src/runtime/scheduler.cpp deleted file mode 100644 index 9bec297e..00000000 --- a/src/runtime/scheduler.cpp +++ /dev/null @@ -1,254 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/runtime/scheduler.hpp" -#include "kmm/runtime/task.hpp" - -namespace kmm { - -static constexpr size_t NUM_DEFAULT_QUEUES = 3; -static constexpr size_t QUEUE_JOIN = 0; -static constexpr size_t QUEUE_MISC = 1; -static constexpr size_t QUEUE_HOST = 2; -static constexpr size_t QUEUE_DEVICES = 3; - -struct QueueSlot { - std::shared_ptr inner; -}; - -bool operator<(const QueueSlot& lhs, const QueueSlot& rhs) { - return lhs.inner->id() > rhs.inner->id(); -} - -struct SchedulerQueue { - size_t max_concurrent_jobs = std::numeric_limits::max(); - size_t num_jobs_active = 0; - std::priority_queue tasks; - - void push_job(const TaskRecord* predecessor, std::shared_ptr record); - std::shared_ptr pop_job(); - void scheduled_job(const TaskRecord& record); - void completed_job(const TaskRecord& record); -}; - -void SchedulerQueue::push_job(const TaskRecord* predecessor, std::shared_ptr record) { - this->tasks.push(QueueSlot {record}); -} - -std::shared_ptr SchedulerQueue::pop_job() { - if (num_jobs_active >= max_concurrent_jobs) { - return nullptr; - } - - if (tasks.empty()) { - return nullptr; - } - - num_jobs_active++; - auto result = std::move(tasks.top()).inner; - tasks.pop(); - return result; -} - -void SchedulerQueue::scheduled_job(const TaskRecord& record) { - // Nothing to do after scheduling -} - -void SchedulerQueue::completed_job(const TaskRecord& record) { - num_jobs_active--; -} - -Scheduler::Scheduler( - std::shared_ptr device_resources, - std::shared_ptr stream_manager, - std::shared_ptr buffer_registry, - bool debug_mode -) : - m_device_resources(device_resources), - m_stream_manager(stream_manager), - m_buffer_registry(buffer_registry), - m_debug_mode(debug_mode) { - size_t num_devices = m_device_resources->num_contexts(); - m_ready_queues.resize(NUM_DEFAULT_QUEUES + num_devices); - - for (size_t i = 0; i < num_devices; i++) { - m_ready_queues[QUEUE_DEVICES + i].max_concurrent_jobs = 5; - } -} - -Scheduler::~Scheduler() {} - -void Scheduler::submit(EventId event_id, std::unique_ptr task, EventList dependencies) { - auto record = std::make_shared(event_id, std::move(task)); - - spdlog::debug( - "submit task {} (command={}, dependencies={})", - event_id, - record->task->name(), - dependencies - ); - - size_t num_pending = 0; - DeviceEventSet dependency_events; - - for (EventId dep_id : dependencies) { - auto it = m_tasks.find(dep_id); - - if (it == m_tasks.end()) { - continue; - } - - auto& dep = it->second; - dep->successors.push_back(record); - - if (dep->status == TaskRecord::Status::WaitingForCompletion) { - dependency_events.insert(dep->output_events); - } else { - num_pending++; - } - } - - record->status = TaskRecord::Status::AwaitingDependencies; - record->predecessors = std::move(dependencies); - record->queue = &m_ready_queues[determine_queue_id(*record->task)]; - record->predecessors_pending = num_pending; - record->input_events = std::move(dependency_events); - enqueue_if_ready(nullptr, record); - - m_tasks.emplace(event_id, std::move(record)); -} - -bool Scheduler::is_completed(EventId event_id) const { - return m_tasks.find(event_id) == m_tasks.end(); -} - -bool Scheduler::is_idle() const { - return m_running_head == nullptr && m_tasks.empty(); -} - -void Scheduler::make_progress() { - TaskRecord* prev = nullptr; - std::shared_ptr* current_ptr = &m_running_head; - - while (auto current = *current_ptr) { - // In debug mode, only poll the head task (prev == nullptr) - bool should_poll = !m_debug_mode || prev == nullptr; - - if (should_poll && this->poll_completion(*current) == Poll::Ready) { - *current_ptr = std::move(current->next); - } else { - prev = current.get(); - current_ptr = ¤t->next; - } - } - - m_running_tail = prev; - - while (auto record = dequeue_ready_task()) { - start_task(std::move(record)); - } -} - -size_t Scheduler::determine_queue_id(const Task& task) { - if (dynamic_cast(&task) != nullptr) { - return QUEUE_JOIN; - } else if (dynamic_cast(&task) != nullptr) { - return QUEUE_HOST; - } else if (const auto* p = dynamic_cast(&task)) { - return QUEUE_DEVICES + p->resource_id().as_device(); - } else { - return QUEUE_MISC; - } -} - -void Scheduler::enqueue_if_ready( - const TaskRecord* predecessor, - const std::shared_ptr& task -) { - if (task->status != TaskRecord::Status::AwaitingDependencies) { - return; - } - - if (task->predecessors_pending > 0) { - return; - } - - task->status = TaskRecord::Status::ReadyToStart; - task->queue->push_job(predecessor, task); -} - -std::shared_ptr Scheduler::dequeue_ready_task() { - for (auto& q : m_ready_queues) { - if (auto result = q.pop_job()) { - return result; - } - } - - return nullptr; -} - -void Scheduler::start_task(std::shared_ptr record) { - KMM_ASSERT(record->status == TaskRecord::Status::ReadyToStart); - - if (poll_completion(*record) == Poll::Pending) { - if (auto* old_tail = std::exchange(m_running_tail, record.get())) { - old_tail->next = std::move(record); - } else { - m_running_head = std::move(record); - } - } -} - -Poll Scheduler::poll_completion(TaskRecord& record) { - if (record.status == TaskRecord::Status::ReadyToStart) { - spdlog::debug( - "scheduling task {} (command={}, GPU deps={})", - record.id(), - record.task->name(), - record.input_events - ); - - record.status = TaskRecord::Status::Running; - record.task->start(record.input_events); - } - - if (record.status == TaskRecord::Status::Running) { - if (record.task->poll(record, *this, record.output_events) != Poll::Ready) { - return Poll::Pending; - } - - spdlog::debug( - "scheduled task {} (command={}, GPU event={})", - record.id(), - record.task->name(), - record.output_events - ); - - record.queue->scheduled_job(record); - record.status = TaskRecord::Status::WaitingForCompletion; - - for (const auto& succ : record.successors) { - succ->input_events.insert(record.output_events); - succ->predecessors_pending -= 1; - enqueue_if_ready(&record, succ); - } - } - - if (record.status == TaskRecord::Status::WaitingForCompletion) { - if (!m_stream_manager->is_ready(record.output_events)) { - return Poll::Pending; - } - - spdlog::debug("completed task {} (command={})", record.id(), record.task->name()); - m_tasks.erase(record.event_id); - record.queue->completed_job(record); - record.status = TaskRecord::Status::Completed; - record.task = nullptr; - } - - KMM_ASSERT(record.status == TaskRecord::Status::Completed); - return Poll::Ready; -} - -} // namespace kmm diff --git a/src/runtime/stream_manager.cpp b/src/runtime/stream_manager.cpp deleted file mode 100644 index 4688e4db..00000000 --- a/src/runtime/stream_manager.cpp +++ /dev/null @@ -1,583 +0,0 @@ -#include -#include -#include - -#include "spdlog/spdlog.h" - -#include "kmm/runtime/stream_manager.hpp" - -namespace kmm { - -using Callback = std::pair; - -struct CompareCallback { - bool operator()(const Callback& a, const Callback& b) const { - return a.first > b.first; - } -}; - -struct DeviceStreamManager::StreamState { - // GCC 9.4 does not allow noexcept in move constructor when using a std::vector. - // This is why we explicitly define them as not being noexcept. - StreamState(const StreamState&) = delete; - StreamState& operator=(const StreamState&) = delete; - StreamState(StreamState&&) /*noexcept*/ = default; - StreamState& operator=(StreamState&&) /*noexcept*/ = default; - - public: - StreamState(size_t pool_index, GPUContextHandle c, g_stream_t s, bool delete_stream_on_exit) : - pool_index(pool_index), - context(c), - gpu_stream(s), - delete_stream_on_exit(delete_stream_on_exit) {} - - size_t pool_index; - GPUContextHandle context; - g_stream_t gpu_stream; - bool delete_stream_on_exit = true; - std::deque pending_events; - DeviceEvent::index_type first_pending_index = 1; - std::priority_queue, CompareCallback> callbacks_heap; -}; - -struct DeviceStreamManager::EventPool { - KMM_NOT_COPYABLE(EventPool) - - public: - EventPool(GPUContextHandle context) : m_context(context) {} - ~EventPool(); - g_event_t pop(); - void push(g_event_t event); - - GPUContextHandle m_context; - std::vector m_events; -}; - -DeviceStreamManager::DeviceStreamManager() {} - -DeviceStreamManager::~DeviceStreamManager() { - for (auto& stream : m_streams) { - GPUContextGuard guard {stream.context}; - - for (const auto& gpu_event : stream.pending_events) { - KMM_GPU_CHECK(g_event_synchronize(gpu_event)); - KMM_ASSERT(g_event_synchronize(gpu_event) == G_SUCCESS); - - stream.first_pending_index += 1; - m_event_pools[stream.pool_index].push(gpu_event); - } - - KMM_GPU_CHECK(g_stream_synchronize(stream.gpu_stream)); - KMM_ASSERT(g_stream_query(stream.gpu_stream) == G_SUCCESS); - - if (stream.delete_stream_on_exit) { - KMM_GPU_CHECK(g_stream_destroy(stream.gpu_stream)); - } - } -} - -auto find_pool_for_context( - GPUContextHandle context, - std::vector& m_event_pools -) { - bool found_pool = false; - size_t pool_index; - - for (size_t i = 0; i < m_event_pools.size(); i++) { - if (m_event_pools[i].m_context == context) { - found_pool = true; - pool_index = i; - } - } - - if (!found_pool) { - pool_index = m_event_pools.size(); - m_event_pools.push_back(context); - } - - return pool_index; -} - -DeviceStream DeviceStreamManager::create_stream(GPUContextHandle context, bool high_priority) { - GPUContextGuard guard {context}; - - int least_priority; - int greatest_priority; - KMM_GPU_CHECK(g_ctx_get_stream_priority_range(&least_priority, &greatest_priority)); - int priority = high_priority ? greatest_priority : least_priority; - - size_t index = m_streams.size(); - g_stream_t gpu_stream; - KMM_GPU_CHECK(g_stream_create_with_priority(&gpu_stream, G_STREAM_NON_BLOCKING, priority)); - - size_t pool_index = find_pool_for_context(context, m_event_pools); - m_streams.emplace_back(pool_index, context, gpu_stream, true); - - return checked_cast(index); -} - -DeviceStream DeviceStreamManager::get_or_add_stream( - GPUContextHandle context, - g_stream_t gpu_stream -) { - for (size_t i = 0; i < m_streams.size(); i++) { - auto& stream = m_streams[i]; - - if (stream.gpu_stream == gpu_stream) { - KMM_ASSERT(stream.context == context); - return checked_cast(i); - } - } - - size_t index = m_streams.size(); - size_t pool_index = find_pool_for_context(context, m_event_pools); - m_streams.emplace_back(pool_index, context, gpu_stream, false); - - return checked_cast(index); -} - -void DeviceStreamManager::wait_until_idle() const { - for (const auto& stream : m_streams) { - KMM_GPU_CHECK(g_stream_synchronize(stream.gpu_stream)); - } -} - -void DeviceStreamManager::wait_until_ready(DeviceStream stream) const { - KMM_GPU_CHECK(g_stream_synchronize(get(stream))); -} - -void DeviceStreamManager::wait_until_ready(DeviceEvent event) const { - KMM_ASSERT(event.stream() < m_streams.size()); - const auto& src_stream = m_streams[event.stream()]; - - if (event.index() < src_stream.first_pending_index) { - return; - } - - auto offset = event.index() - src_stream.first_pending_index; - g_event_t gpu_event = src_stream.pending_events.at(offset); - - GPUContextGuard guard {src_stream.context}; - KMM_GPU_CHECK(g_event_synchronize(gpu_event)); -} - -void DeviceStreamManager::wait_until_ready(const DeviceEventSet& events) const { - for (DeviceEvent e : events) { - wait_until_ready(e); - } -} - -bool DeviceStreamManager::is_idle() const { - for (const auto& stream : m_streams) { - if (!stream.pending_events.empty()) { - return false; - } - - if (!stream.callbacks_heap.empty()) { - return false; - } - } - - for (const auto& stream : m_streams) { - GPUContextGuard guard {stream.context}; - KMM_GPU_CHECK(g_stream_synchronize(stream.gpu_stream)); - KMM_GPU_CHECK(g_stream_synchronize(nullptr)); - } - - return true; -} - -bool DeviceStreamManager::is_ready(DeviceStream stream) const noexcept { - KMM_ASSERT(stream < m_streams.size()); - return m_streams[stream].pending_events.empty(); -} - -bool DeviceStreamManager::is_ready(DeviceEvent event) const noexcept { - KMM_ASSERT(event.stream() < m_streams.size()); - return m_streams[event.stream()].first_pending_index > event.index(); -} - -bool DeviceStreamManager::is_ready(const DeviceEventSet& events) const noexcept { - for (DeviceEvent e : events) { - if (!is_ready(e)) { - return false; - } - } - - return true; -} - -bool DeviceStreamManager::is_ready(DeviceEventSet& events) const noexcept { - return events.remove_ready_trailing(*this); -} - -void DeviceStreamManager::attach_callback(DeviceEvent event, NotifyHandle callback) { - KMM_ASSERT(event.stream() < m_streams.size()); - auto& stream = m_streams[event.stream()]; - stream.callbacks_heap.emplace(event.index(), std::move(callback)); -} - -void DeviceStreamManager::attach_callback(DeviceStream stream, NotifyHandle callback) { - attach_callback(record_event(stream), std::move(callback)); -} - -DeviceEvent DeviceStreamManager::record_event(DeviceStream stream_id) { - KMM_ASSERT(stream_id < m_streams.size()); - auto& stream = m_streams[stream_id]; - - auto event_index = stream.first_pending_index + stream.pending_events.size(); - auto event = DeviceEvent {stream_id, event_index}; - - g_event_t gpu_event = m_event_pools[stream.pool_index].pop(); - stream.pending_events.push_back(gpu_event); - - KMM_GPU_CHECK(g_event_record(gpu_event, stream.gpu_stream)); - - spdlog::trace("GPU stream {} records new GPU event {}", stream_id, event); - return event; -} - -void DeviceStreamManager::wait_on_default_stream(DeviceStream stream_id) { - KMM_ASSERT(stream_id < m_streams.size()); - auto& stream = m_streams[stream_id]; - - g_event_t gpu_event = m_event_pools[stream.pool_index].pop(); - m_event_pools[stream.pool_index].push(gpu_event); - - KMM_GPU_CHECK(g_event_record(gpu_event, 0)); - KMM_GPU_CHECK(g_stream_wait_event(stream.gpu_stream, gpu_event, G_EVENT_WAIT_DEFAULT)); -} - -void DeviceStreamManager::wait_for_event(DeviceStream stream, DeviceEvent event) const { - KMM_ASSERT(event.stream() < m_streams.size()); - KMM_ASSERT(stream < m_streams.size()); - - // Stream never needs to wait on events from itself - if (event.stream() == stream) { - return; - } - - const auto& src_stream = m_streams.at(event.stream()); - const auto& dst_stream = m_streams.at(stream); - - // Event has already completed, no need to wait. - if (event.index() < src_stream.first_pending_index) { - return; - } - - auto offset = event.index() - src_stream.first_pending_index; - g_event_t gpu_event = src_stream.pending_events.at(offset); - KMM_GPU_CHECK(g_stream_wait_event(dst_stream.gpu_stream, gpu_event, G_EVENT_WAIT_DEFAULT)); - - spdlog::trace("GPU stream {} must wait on GPU event {}", stream, event); -} - -void DeviceStreamManager::wait_for_events( - DeviceStream stream, - const DeviceEvent* begin, - const DeviceEvent* end -) const { - for (const auto* it = begin; it != end; it++) { - wait_for_event(stream, *it); - } -} - -void DeviceStreamManager::wait_for_events(DeviceStream stream, const DeviceEventSet& events) const { - wait_for_events(stream, events.begin(), events.end()); -} - -void DeviceStreamManager::wait_for_events( - DeviceStream stream, - const std::vector& events -) const { - wait_for_events(stream, &*events.begin(), &*events.end()); -} - -bool DeviceStreamManager::event_happens_before(DeviceEvent source, DeviceEvent target) { - return source.stream() == target.stream() && source.index() < target.index(); -} - -GPUContextHandle DeviceStreamManager::context(DeviceStream stream) const { - KMM_ASSERT(stream.get() < m_streams.size()); - return m_streams[stream.get()].context; -} - -g_stream_t DeviceStreamManager::get(DeviceStream stream) const { - KMM_ASSERT(stream < m_streams.size()); - return m_streams[stream].gpu_stream; -} - -bool DeviceStreamManager::make_progress() { - bool update_happened = false; - - for (size_t i = 0; i < m_streams.size(); i++) { - if (make_progress_for_stream(static_cast(i))) { - update_happened = true; - } - } - - return update_happened; -} - -bool DeviceStreamManager::make_progress_for_stream(DeviceStream stream_index) { - auto update_happened = false; - auto& stream = m_streams[stream_index]; - - if (!stream.pending_events.empty()) { - GPUContextGuard guard {stream.context}; - - do { - g_event_t gpu_event = stream.pending_events[0]; - g_result_t result = g_event_query(gpu_event); - - if (result == G_ERROR_NOT_READY) { - break; - } - - if (result != G_SUCCESS) { - throw GPUDriverException("`gpuEventQuery` failed", result); - } - - spdlog::trace( - "GPU event {} completed", - DeviceEvent(stream_index, stream.first_pending_index) - ); - - stream.first_pending_index += 1; - stream.pending_events.pop_front(); - m_event_pools[stream.pool_index].push(gpu_event); - update_happened = true; - } while (!stream.pending_events.empty()); - } - - while (!stream.callbacks_heap.empty()) { - const auto& [event_index, handle] = stream.callbacks_heap.top(); - - if (event_index >= stream.first_pending_index) { - break; - } - - handle.notify(); - stream.callbacks_heap.pop(); - update_happened = true; - } - - return update_happened; -} - -DeviceStreamManager::EventPool::~EventPool() { - GPUContextGuard guard {m_context}; - - for (const auto& gpu_event : m_events) { - KMM_GPU_CHECK(g_event_destroy(gpu_event)); - } -} - -g_event_t DeviceStreamManager::EventPool::pop() { - g_event_t gpu_event; - - if (m_events.empty()) { - GPUContextGuard guard {m_context}; - KMM_GPU_CHECK(g_event_create(&gpu_event, G_EVENT_DISABLE_TIMING)); - } else { - gpu_event = m_events.back(); - m_events.pop_back(); - } - - return gpu_event; -} - -void DeviceStreamManager::EventPool::push(g_event_t event) { - m_events.push_back(event); -} - -DeviceEventSet::DeviceEventSet(DeviceEvent e) { - m_events.push_back(e); -} - -DeviceEventSet::DeviceEventSet(std::initializer_list e) { - m_events.insert_all(e.begin(), e.end()); -} - -DeviceEventSet& DeviceEventSet::operator=(std::initializer_list e) { - clear(); - m_events.insert_all(e.begin(), e.end()); - return *this; -} - -void DeviceEventSet::insert(DeviceEvent e) noexcept { - static constexpr size_t INVALID_INDEX = std::numeric_limits::max(); - size_t found_index = INVALID_INDEX; - - if (e.is_null()) { - return; - } - - for (size_t i = 0; i < m_events.size(); i++) { - if (m_events[i].stream() == e.stream()) { - found_index = i; - } - } - - if (found_index != INVALID_INDEX) { - m_events[found_index] = std::max(m_events[found_index], e); - } else { - KMM_ASSERT(m_events.try_push_back(e)); - } -} - -void DeviceEventSet::insert(const DeviceEventSet& that) noexcept { - static constexpr size_t INVALID_INDEX = std::numeric_limits::max(); - size_t num_old_events = m_events.size(); - - for (auto e : that.m_events) { - size_t found_index = INVALID_INDEX; - - for (size_t i = 0; i < num_old_events; i++) { - if (m_events[i].stream() == e.stream()) { - found_index = i; - } - } - - if (found_index != INVALID_INDEX) { - m_events[found_index] = std::max(m_events[found_index], e); - } else { - KMM_ASSERT(m_events.try_push_back(e)); - } - } -} - -void DeviceEventSet::insert(DeviceEventSet&& that) noexcept { - if (that.m_events.size() > this->m_events.size()) { - std::swap(this->m_events, that.m_events); - } - - insert(that); -} - -bool DeviceEventSet::remove_ready(const DeviceStreamManager& m) noexcept { - size_t old_size = m_events.size(); - size_t new_size = 0; - - for (size_t index = 0; index < old_size; index++) { - auto event = m_events[index]; - - if (!m.is_ready(event)) { - m_events[new_size] = event; - new_size++; - } - } - - m_events.truncate(new_size); - return new_size > 0; -} - -bool DeviceEventSet::remove_ready_trailing(const DeviceStreamManager& m) noexcept { - for (size_t new_size = m_events.size(); new_size > 0; new_size--) { - auto event = m_events[new_size - 1]; - - // event is not ready, truncate events up to `new_size`. - if (!m.is_ready(event)) { - m_events.truncate(new_size); - return false; - } - } - - // all events are ready - m_events.clear(); - return true; -} - -DeviceEventSet DeviceEventSet::extract_events_for_context( - const DeviceStreamManager& manager, - GPUContextHandle context -) { - // Remove all events that have completed. - remove_ready(manager); - - // Push all events with a different context to the front of the list. - auto* mid = std::partition(m_events.begin(), m_events.end(), [&](DeviceEvent e) { - return manager.context(e.stream()) != context; - }); - - // If all events have the same context, then we can just return the current set. - if (m_events.begin() == mid) { - return std::move(*this); - } - - // If all events have a different context, then we can just return an empty set. - if (m_events.end() == mid) { - return DeviceEventSet {}; - } - - DeviceEventSet result; - result.m_events.insert_all(mid, m_events.end()); - m_events.truncate(static_cast(mid - m_events.begin())); - return result; -} - -void DeviceEventSet::clear() noexcept { - m_events.clear(); -} - -bool DeviceEventSet::is_empty() const noexcept { - return m_events.is_empty(); -} - -const DeviceEvent* DeviceEventSet::begin() const noexcept { - return m_events.begin(); -} - -const DeviceEvent* DeviceEventSet::end() const noexcept { - return m_events.end(); -} - -DeviceEventSet operator|(const DeviceEventSet& a, const DeviceEventSet& b) noexcept { - DeviceEventSet result = a; - result.insert(b); - return result; -} - -std::ostream& operator<<(std::ostream& f, const DeviceStream& e) { - return f << uint32_t(e.get()); -} - -std::ostream& operator<<(std::ostream& f, const DeviceEvent& e) { - if (e.m_event_and_stream == 0) { - return f << ""; - } - - return f << e.stream() << ":" << e.index(); -} - -std::ostream& operator<<(std::ostream& f, const DeviceEventSet& events) { - // Sort events - auto sorted_events = std::vector {events.begin(), events.end()}; - std::sort(sorted_events.begin(), sorted_events.end()); - - // Remove duplicates - auto it = std::unique(sorted_events.begin(), sorted_events.end()); - sorted_events.erase(it, sorted_events.end()); - - bool is_first = true; - f << "["; - - for (auto e : sorted_events) { - // Skip empty events - if (e.is_null()) { - continue; - } - - if (!is_first) { - f << ", "; - } - - is_first = false; - f << e; - } - - f << "]"; - return f; -} - -} // namespace kmm diff --git a/src/runtime/system_info.cpp b/src/runtime/system_info.cpp new file mode 100644 index 00000000..1b65027c --- /dev/null +++ b/src/runtime/system_info.cpp @@ -0,0 +1,150 @@ +#include +#include +#include + +#include "kmm/runtime/system_info.hpp" +#include "kmm/utils/gpu_utils.hpp" + +namespace kmm { + +DeviceInfo::DeviceInfo(DeviceId id, g_context_t context) : + m_id(id), + m_context(context), + m_context_id(context) { + GPUContextGuard guard {context}; + + KMM_GPU_CHECK(g_ctx_get_device(&m_device)); + + std::array name_buf {}; + KMM_GPU_CHECK(g_device_get_name(name_buf.data(), static_cast(name_buf.size()), m_device)); + m_name = name_buf.data(); + + KMM_GPU_CHECK(g_device_total_mem(&m_total_memory, m_device)); + m_memory_capacity = m_total_memory; + + KMM_GPU_CHECK(g_device_get_attribute( + &m_compute_capability_major, + G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR, + m_device + )); + KMM_GPU_CHECK(g_device_get_attribute( + &m_compute_capability_minor, + G_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR, + m_device + )); +} + +int DeviceInfo::attribute(g_device_attribute_t attrib) const { + KMM_ASSERT(static_cast(attrib) < NUM_ATTRIBUTES); + + int value; + KMM_GPU_CHECK(g_device_get_attribute(&value, attrib, m_device)); + return value; +} + +int DeviceInfo::max_threads_per_block() const { + return attribute(G_DEVICE_ATTRIBUTE_MAX_THREADS_PER_BLOCK); +} + +dim3 DeviceInfo::max_block_dim() const { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + return dim3( + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_X)), + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Y)), + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Z)) + ); +#else + throw std::runtime_error("unsupported operation"); +#endif +} + +dim3 DeviceInfo::max_grid_dim() const { +#if defined(KMM_USE_CUDA) || defined(KMM_USE_HIP) + return dim3( + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_X)), + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Y)), + static_cast(attribute(G_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z)) + ); +#else + throw std::runtime_error("unsupported operation"); +#endif +} + +std::pair DeviceInfo::compute_capability() const { + return {m_compute_capability_major, m_compute_capability_minor}; +} + +static std::vector query_all_devices() { + KMM_GPU_CHECK(g_init(0)); + + int count = 0; + KMM_GPU_CHECK(g_device_get_count(&count)); + + size_t n = static_cast(count) < MAX_DEVICES ? static_cast(count) : MAX_DEVICES; + + std::vector devices; + devices.reserve(n); + + for (size_t i = 0; i < n; i++) { + g_device_t ordinal; + KMM_GPU_CHECK(g_device_get(&ordinal, static_cast(i))); + + g_context_t context; + KMM_GPU_CHECK(g_device_primary_ctx_retain(&context, ordinal)); + + devices.emplace_back(DeviceId(i), context); + } + + return devices; +} + +SystemInfo::SystemInfo() : SystemInfo(query_all_devices()) {} + +SystemInfo::SystemInfo(std::vector devices) : m_devices(std::move(devices)) {} + +SystemInfo::~SystemInfo() { + // Never throw out of a destructor: ignore release failures rather than routing them through + // KMM_GPU_CHECK. + for (const auto& info : m_devices) { + g_device_primary_ctx_release(info.device_ordinal()); + } +} + +size_t SystemInfo::num_devices() const { + return m_devices.size(); +} + +const DeviceInfo& SystemInfo::device(DeviceId id) const { + KMM_ASSERT(id.get() < m_devices.size()); + return m_devices[id.get()]; +} + +const DeviceInfo& SystemInfo::device_by_ordinal(g_device_t ordinal) const { + for (const auto& info : m_devices) { + if (info.device_ordinal() == ordinal) { + return info; + } + } + + throw std::runtime_error("no device found with the given device ordinal"); +} + +MemoryId SystemInfo::affinity_memory(DeviceId device_id) const { + return device(device_id).memory_id(); +} + +const DeviceInfo& SystemInfo::device_from_context(g_context_t context) const { + for (const auto& info : m_devices) { + if (info.context() == context) { + return info; + } + } + + throw std::runtime_error("no device found with the given CUDA context"); +} + +const DeviceInfo& SystemInfo::device_from_stream(g_stream_t stream) const { + return device_from_context(context_from_stream(stream)); +} + +} // namespace kmm diff --git a/src/runtime/task.cpp b/src/runtime/task.cpp deleted file mode 100644 index 5a8d735c..00000000 --- a/src/runtime/task.cpp +++ /dev/null @@ -1,364 +0,0 @@ -#include "kmm/memops/gpu_copy.hpp" -#include "kmm/memops/gpu_fill.hpp" -#include "kmm/memops/gpu_reduction.hpp" -#include "kmm/memops/host_copy.hpp" -#include "kmm/memops/host_fill.hpp" -#include "kmm/memops/host_reduction.hpp" -#include "kmm/runtime/scheduler.hpp" -#include "kmm/runtime/task.hpp" - -namespace kmm { - -static PoisonException make_poison_exception(TaskRecord& record, const std::exception& error) { - if (const auto* reason = dynamic_cast(&error)) { - return *reason; - } - - return fmt::format("task {} failed due to error: {}", record.id(), error.what()); -} - -void JoinTask::start(const DeviceEventSet& input_events) { - m_dependencies = input_events; -} - -Poll JoinTask::poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) { - output_events.insert(std::move(m_dependencies)); - return Poll::Ready; -} - -void DeleteBufferTask::start(const DeviceEventSet& input_events) { - m_dependencies = input_events; -} - -Poll DeleteBufferTask::poll( - TaskRecord& record, - Scheduler& scheduler, - DeviceEventSet& output_events -) { - scheduler.buffers().remove(m_buffer_id); - output_events.insert(m_dependencies); - return Poll::Ready; -} - -void HostTask::start(const DeviceEventSet& input_events) { - KMM_ASSERT(m_status == Status::Init); - m_dependencies = input_events; - m_status = Status::CreateBuffers; -} - -Poll HostTask::poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) { - if (m_status == Status::CreateBuffers) { - try { - m_requests = scheduler.buffers().create_requests(m_buffers); - m_status = Status::PollingBuffers; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::PollingBuffers) { - try { - if (scheduler.buffers().poll_requests(m_requests, m_dependencies) == Poll::Pending) { - return Poll::Pending; - } - - m_status = Status::PollingDependencies; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::PollingDependencies) { - try { - if (!scheduler.streams().is_ready(m_dependencies)) { - return Poll::Pending; - } - - m_future = submit(scheduler, scheduler.buffers().access_requests(m_requests)); - m_status = Status::Running; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::Running) { - try { - if (m_future.wait_for(std::chrono::seconds(0)) == std::future_status::timeout) { - return Poll::Pending; - } - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - } - - m_status = Status::Completing; - } - - if (m_status == Status::Completing) { - scheduler.buffers().release_requests(m_requests); - m_status = Status::Completed; - } - - return Poll::Ready; -} - -std::future ExecuteHostTask::submit( - Scheduler& scheduler, - std::vector accessors -) { - KMM_ASSERT(m_task != nullptr); - auto* task = m_task.get(); - - return std::async(std::launch::async, [=] { - auto host = HostResource {}; - auto context = TaskContext {std::move(accessors)}; - task->execute(host, context); - }); -} - -std::future CopyHostTask::submit( - Scheduler& scheduler, - std::vector accessors -) { - KMM_ASSERT(accessors[0].layout.size_in_bytes >= m_copy.minimum_source_bytes_needed()); - KMM_ASSERT(accessors[1].layout.size_in_bytes >= m_copy.minimum_destination_bytes_needed()); - KMM_ASSERT(accessors[1].is_writable); - - return std::async(std::launch::async, [=] { - execute_copy(accessors[0].address, accessors[1].address, m_copy); - }); -} - -std::future ReductionHostTask::submit( - Scheduler& scheduler, - std::vector accessors -) { - return std::async(std::launch::async, [=] { - execute_reduction(accessors[0].address, accessors[1].address, m_reduction); - }); -} - -std::future FillHostTask::submit( - Scheduler& scheduler, - std::vector accessors -) { - return std::async(std::launch::async, [=] { execute_fill(accessors[0].address, m_fill); }); -} - -void DeviceTask::start(const DeviceEventSet& input_events) { - KMM_ASSERT(m_status == Status::Init); - m_dependencies = input_events; - m_status = Status::CreateBuffers; -} - -Poll DeviceTask::poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) { - if (m_status == Status::CreateBuffers) { - try { - m_requests = scheduler.buffers().create_requests(m_buffers); - m_status = Status::PollingBuffers; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::PollingBuffers) { - try { - if (scheduler.buffers().poll_requests(m_requests, m_dependencies) == Poll::Pending) { - return Poll::Pending; - } - - // Remove the `local` events from the list of dependencies. These events - // have the same context as the current device, and thus can be directly put - // as dependencies on the current device stream. - m_local_dependencies = m_dependencies.extract_events_for_context( - scheduler.streams(), - scheduler.devices().context(m_resource.as_device()) - ); - - output_events.insert(m_local_dependencies); - m_status = Status::PollingDependencies; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::PollingDependencies) { - if (!scheduler.streams().is_ready(m_dependencies)) { - return Poll::Pending; - } - - m_status = Status::Running; - } - - if (m_status == Status::Running) { - try { - m_execution_event = scheduler.devices().submit( - m_resource.as_device(), - m_resource.stream_affinity(), - m_local_dependencies, - *this, - scheduler.buffers().access_requests(m_requests) - ); - - output_events.insert(m_execution_event); - m_status = Status::Completing; - } catch (const std::exception& e) { - scheduler.buffers().poison_all(m_buffers, make_poison_exception(record, e)); - m_status = Status::Completing; - } - } - - if (m_status == Status::Completing) { - scheduler.buffers().release_requests(m_requests, m_execution_event); - m_status = Status::Completed; - } - - return Poll::Ready; -} - -void ExecuteDeviceTask::execute(DeviceResource& device, std::vector accessors) { - KMM_ASSERT(m_task != nullptr); - auto context = TaskContext {std::move(accessors)}; - m_task->execute(device, context); -} - -void CopyDeviceTask::execute(DeviceResource& device, std::vector accessors) { - KMM_ASSERT(accessors[0].layout.size_in_bytes >= m_copy.minimum_source_bytes_needed()); - KMM_ASSERT(accessors[1].layout.size_in_bytes >= m_copy.minimum_destination_bytes_needed()); - KMM_ASSERT(accessors[1].is_writable); - - execute_gpu_d2d_copy_async( - device, - reinterpret_cast(accessors[0].address), - reinterpret_cast(accessors[1].address), - m_copy - ); -} - -void ReductionDeviceTask::execute(DeviceResource& device, std::vector accessors) { - execute_gpu_reduction_async( - device, - reinterpret_cast(accessors[0].address), - reinterpret_cast(accessors[1].address), - m_reduction - ); -} - -void FillDeviceTask::execute(DeviceResource& device, std::vector accessors) { - execute_gpu_fill_async(device, reinterpret_cast(accessors[0].address), m_fill); -} - -void PrefetchTask::start(const DeviceEventSet& input_events) { - m_dependencies = input_events; -} - -Poll PrefetchTask::poll(TaskRecord& record, Scheduler& scheduler, DeviceEventSet& output_events) { - if (m_status == Status::Init) { - m_requests = scheduler.buffers().create_requests(m_buffers); - m_status = Status::Polling; - } - - if (m_status == Status::Polling) { - if (scheduler.buffers().poll_requests(m_requests, m_dependencies) == Poll::Pending) { - return Poll::Pending; - } - - scheduler.buffers().release_requests(m_requests); - m_status = Status::Completing; - } - - if (m_status == Status::Completing) { - if (!scheduler.streams().is_ready(m_dependencies)) { - return Poll::Pending; - } - - m_status = Status::Completed; - } - - return Poll::Ready; -} - -std::unique_ptr build_task_for_command(Command&& command) { - if (std::get_if(&command) != nullptr) { - return std::make_unique(); - - } else if (const auto* e = std::get_if(&command)) { - return std::make_unique(e->id); - - } else if (const auto* e = std::get_if(&command)) { - return std::make_unique(e->buffer_id, e->memory_id); - - } else if (auto* e = std::get_if(&command)) { - auto proc = e->processor_id; - - if (proc.is_device()) { - return std::make_unique(proc, std::move(e->task), e->buffers); - } else { - return std::make_unique(std::move(e->task), e->buffers); - } - - } else if (const auto* e = std::get_if(&command)) { - auto src_mem = e->src_memory; - auto dst_mem = e->dst_memory; - - if (src_mem.is_host() && dst_mem.is_host()) { - return std::make_unique(e->src_buffer, e->dst_buffer, e->definition); - } else if (dst_mem.is_device()) { - return std::make_unique( - dst_mem.as_device(), - e->src_buffer, - e->dst_buffer, - e->definition - ); - } else if (src_mem.is_device()) { - return std::make_unique( - src_mem.as_device(), - e->src_buffer, - e->dst_buffer, - e->definition - ); - } else { - KMM_PANIC("unsupported copy"); - } - - } else if (const auto* e = std::get_if(&command)) { - auto memory_id = e->memory_id; - - if (memory_id.is_device()) { - return std::make_unique( - memory_id.as_device(), - e->src_buffer, - e->dst_buffer, - std::move(e->definition) - ); - } else { - return std::make_unique( - e->src_buffer, - e->dst_buffer, - std::move(e->definition) - ); - } - - } else if (const auto* e = std::get_if(&command)) { - auto memory_id = e->memory_id; - - if (memory_id.is_device()) { - return std::make_unique( - memory_id.as_device(), - e->dst_buffer, - std::move(e->definition) - ); - } else { - return std::make_unique(e->dst_buffer, std::move(e->definition)); - } - - } else { - KMM_PANIC_FMT("could not handle unknown command: {}", command); - } -} - -} // namespace kmm diff --git a/src/runtime/task_graph.cpp b/src/runtime/task_graph.cpp deleted file mode 100644 index c5d9ad0b..00000000 --- a/src/runtime/task_graph.cpp +++ /dev/null @@ -1,84 +0,0 @@ -#include - -#include "spdlog/spdlog.h" - -#include "kmm/runtime/task_graph.hpp" - -namespace kmm { - -EventId TaskGraphState::commit( - TaskGraph& g, - std::vector& staged_nodes, - std::vector>& staged_buffers -) { - KMM_ASSERT(g.m_state == this); - m_last_barrier_id = g.insert_barrier(); - staged_nodes = std::move(g.m_staged_nodes); - staged_buffers = std::move(g.m_staged_buffers); - return m_last_barrier_id; -} - -TaskGraph::TaskGraph(TaskGraphState* state) : m_state(state) {} - -EventId TaskGraph::join_events(EventList deps) { - if (deps.size() == 0) { - return EventId(); - } - - if (std::equal(deps.begin() + 1, deps.end(), deps.begin())) { - return deps[0]; - } - - return insert_node(CommandEmpty {}, std::move(deps)); -} - -BufferId TaskGraph::create_buffer(BufferLayout layout) { - auto buffer_id = m_state->m_next_buffer_id; - m_state->m_next_buffer_id = BufferId(buffer_id + 1); - m_staged_buffers.emplace_back(buffer_id, layout); - return buffer_id; -} - -EventId TaskGraph::delete_buffer(BufferId id, EventList deps) { - return insert_node(CommandBufferDelete {id}, std::move(deps)); -} - -EventId TaskGraph::insert_barrier() { - if (m_events_since_last_barrier.is_empty()) { - return m_state->m_last_barrier_id; - } - - EventList deps = std::move(m_events_since_last_barrier); - deps.push_back(m_state->m_last_barrier_id); - - return join_events(std::move(deps)); -} - -EventId TaskGraph::insert_compute_task( - ResourceId process_id, - std::unique_ptr task, - std::vector buffers, - EventList deps -) { - return insert_node( - CommandExecute { - .processor_id = process_id, - .task = std::move(task), - .buffers = std::move(buffers) - }, - std::move(deps) - ); -} - -EventId TaskGraph::insert_node(Command command, EventList deps) { - auto id = EventId(m_state->m_next_event_id.get()); - m_state->m_next_event_id = EventId(id.get() + 1); - - m_events_since_last_barrier.push_back(id); - m_staged_nodes.push_back( - Node {.id = id, .command = std::move(command), .dependencies = std::move(deps)} - ); - - return id; -} -} // namespace kmm diff --git a/src/backends/cpu.cpp b/src/utils/backends/cpu.cpp similarity index 84% rename from src/backends/cpu.cpp rename to src/utils/backends/cpu.cpp index 2912dc6f..37e63fcb 100644 --- a/src/backends/cpu.cpp +++ b/src/utils/backends/cpu.cpp @@ -1,5 +1,7 @@ -#include "kmm/core/backends.hpp" -#include "kmm/memops/types.hpp" +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/reduction.hpp" +#include "kmm/runtime/memops/types.hpp" +#include "kmm/utils/backends.hpp" namespace kmm { @@ -7,14 +9,26 @@ g_result_t g_ctx_get_device(g_device_t* device) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_ctx_get_id(g_context_t ctx, unsigned long long* ctxId) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_device_get_name(char* name, int len, g_device_t dev) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_device_total_mem(size_t*, g_device_t) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_device_get_attribute(int* value, g_device_attribute_t attribute, g_device_t dev) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_device_can_access_peer(int* canAccessPeer, g_device_t dev, g_device_t peerDev) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_mem_get_info(size_t* free, size_t* total) { return g_result_t(G_ERROR_UNKNOWN); } @@ -158,10 +172,26 @@ g_result_t g_ctx_get_stream_priority_range(int* least, int* greatest) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_stream_create(g_stream_t* stream, unsigned int flags) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_stream_create_with_priority(g_stream_t* stream, unsigned int flags, int priority) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_stream_get_device(g_stream_t hStream, g_device_t* device) { + return g_result_t(G_ERROR_UNKNOWN); +} + +g_result_t g_stream_get_id(g_stream_t hStream, unsigned long long* streamId) { + return g_result_t(G_ERROR_UNKNOWN); +} + +g_result_t g_stream_get_ctx(g_stream_t hStream, g_context_t* pctx) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_stream_query(g_stream_t stream) { return g_result_t(G_ERROR_UNKNOWN); } @@ -275,6 +305,23 @@ g_result_t g_pointer_get_attribute(void* prt, g_pointer_attribute_t attribute, g return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_mem_alloc_managed(g_device_ptr_t* dptr, size_t bytesize, unsigned int flags) { + return g_result_t(G_ERROR_UNKNOWN); +} + +g_result_t g_mem_host_get_device_pointer(g_device_ptr_t* pdptr, void* p, unsigned int Flags) { + return g_result_t(G_ERROR_UNKNOWN); +} + +g_result_t g_mem_prefetch_async( + g_device_ptr_t devPtr, + size_t count, + int device, + g_stream_t hStream +) { + return g_result_t(G_ERROR_UNKNOWN); +} + g_result_t g_ctx_create(g_context_t* ctx, unsigned int flags, g_device_t dev) { return g_result_t(G_ERROR_UNKNOWN); } @@ -299,6 +346,10 @@ g_result_t g_ctx_pop_current(g_context_t* ctx) { return g_result_t(G_ERROR_UNKNOWN); } +g_result_t g_ctx_enable_peer_access(g_context_t peerContext, unsigned int Flags) { + return g_result_t(G_ERROR_UNKNOWN); +} + gpu_error_t gpu_launch_kernel( const void* func, dim3 grid, @@ -374,9 +425,13 @@ void execute_gpu_reduction_async( g_stream_t stream, g_device_ptr_t src_buffer, g_device_ptr_t dst_buffer, - ReductionDef reduction + ReductionDescription reduction ) {} -void execute_gpu_fill_async(g_stream_t stream, g_device_ptr_t dst_buffer, const FillDef& fill) {} +void execute_gpu_fill_async( + g_stream_t stream, + g_device_ptr_t dst_buffer, + const FillDescription& fill +) {} -} // namespace kmm +} // namespace kmm \ No newline at end of file diff --git a/src/utils/gpu_utils.cpp b/src/utils/gpu_utils.cpp index 2ad8b775..f2ffb901 100644 --- a/src/utils/gpu_utils.cpp +++ b/src/utils/gpu_utils.cpp @@ -1,142 +1,93 @@ #include "fmt/format.h" -#include "spdlog/spdlog.h" #include "kmm/utils/gpu_utils.hpp" -#include "kmm/utils/panic.hpp" namespace kmm { void gpu_throw_exception(g_result_t result, const char* file, int line, const char* expression) { - throw GPUDriverException(fmt::format("{} ({}:{})", expression, file, line), result); -} + const char* name = "UNKNOWN_ERROR"; + const char* description = "unknown error"; + g_get_error_name(result, &name); + g_get_error_string(result, &description); -#ifndef KMM_USE_HIP -void gpu_throw_exception(gpu_error_t result, const char* file, int line, const char* expression) { - throw GPURuntimeException(fmt::format("{} ({}:{})", expression, file, line), result); + throw GPUException( + fmt::format("GPU error: {} ({}) at {}:{}: {}", name, description, file, line, expression) + ); } -#endif -void gpu_throw_exception(blas_status_t result, const char* file, int line, const char* expression) { - throw BlasException(fmt::format("{} ({}:{})", expression, file, line), result); +GPUContextGuard::GPUContextGuard(g_context_t context) : m_context(context) { + KMM_GPU_CHECK(g_ctx_push_current(context)); } -GPUDriverException::GPUDriverException(const std::string& message, g_result_t result) : - status(result) { - const char* name = "???"; - const char* description = "???"; - - // Ignore the return code from these functions - g_get_error_name(result, &name); - g_get_error_string(result, &description); - - m_message = fmt::format("GPU driver error: {} ({}): {}", description, name, message); +GPUContextGuard::~GPUContextGuard() { + // Destructors must not throw, so the pop result is discarded rather than checked. + g_context_t popped = nullptr; + g_ctx_pop_current(&popped); } -GPURuntimeException::GPURuntimeException(const std::string& message, gpu_error_t result) : - status(result) { - const char* name = "???"; - const char* description = "???"; - - // Ignore the return code from these functions - name = gpu_get_error_name(result); - description = gpu_get_error_string(result); +GPUContextId::GPUContextId(g_context_t context) { + KMM_GPU_CHECK(g_ctx_get_id(context, &m_id)); +} - m_message = fmt::format("GPU runtime error: {} ({}): {}", description, name, message); +g_context_t context_from_stream(g_stream_t stream) { + g_context_t context; + KMM_GPU_CHECK(g_stream_get_ctx(stream, &context)); + return context; } -BlasException::BlasException(const std::string& message, blas_status_t result) : status(result) { - const char* name = blas_get_status_name(result); - const char* description = blas_get_status_string(result); +GPUStreamId::GPUStreamId(g_stream_t stream) : GPUStreamId(stream, context_from_stream(stream)) {} - m_message = fmt::format("BLAS runtime error: {} ({}): {}", description, name, message); +GPUStreamId::GPUStreamId(g_stream_t stream, g_context_t context) : m_context_id(context) { + KMM_GPU_CHECK(g_stream_get_id(stream, &m_id)); } -GPUContextHandle::GPUContextHandle(g_context_t context, std::shared_ptr lifetime) : - m_context(context), - m_lifetime(std::move(lifetime)) {} - -std::vector get_gpu_devices() { - try { - auto result = g_init(0); - if (result == G_ERROR_NO_DEVICE) { - return {}; - } - - if (result != G_SUCCESS) { - throw GPUDriverException("gpuInit failed", result); - } - - int count = 0; - KMM_GPU_CHECK(g_device_get_count(&count)); - - std::vector devices {}; - for (int i = 0; i < count; i++) { - g_device_t device; - KMM_GPU_CHECK(g_device_get(&device, i)); - devices.push_back(device); - } - - return devices; - } catch (const GPUException& e) { - spdlog::warn("ignored error while initializing: {}", e.what()); - return {}; - } -} +GPUStreamId::GPUStreamId(const GPUStreamRef& stream) : GPUStreamId(stream.stream_id()) {} -std::optional get_gpu_device_by_address(const void* address) { - int ordinal; - g_memory_type_t memory_type; - g_result_t result = g_pointer_get_attribute( - &memory_type, - G_POINTER_ATTRIBUTE_MEMORY_TYPE, - g_device_ptr_t(address) - ); +GPUStreamId::GPUStreamId(const GPUStreamOwner& stream) : + GPUStreamId(static_cast(stream)) {} - if (result == G_SUCCESS && memory_type == G_MEMORYTYPE_DEVICE) { - result = g_pointer_get_attribute( - &ordinal, - G_POINTER_ATTRIBUTE_DEVICE_ORDINAL, - g_device_ptr_t(address) - ); +GPUStreamRef::GPUStreamRef(g_stream_t stream) : + m_context(context_from_stream(stream)), + m_stream(stream), + m_stream_id(stream, m_context) {} - if (result == G_SUCCESS) { - return g_device_t {ordinal}; - } - } +GPUStreamOwner::GPUStreamOwner(g_context_t context, unsigned int flags) : + m_stream([&]() { + g_stream_t result; + GPUContextGuard guard(context); + KMM_GPU_CHECK(g_stream_create(&result, flags)); + return result; + }()) {} - return std::nullopt; +GPUStreamOwner::~GPUStreamOwner() { + destroy(); } -GPUContextHandle GPUContextHandle::create_context_for_device(g_device_t device) { - int flags = G_CTX_MAP_HOST; - g_context_t context; - KMM_GPU_CHECK(g_ctx_create(&context, flags, device)); - - auto lifetime = std::shared_ptr(nullptr, [=](const void* ignore) { - KMM_ASSERT(g_ctx_destroy(context) == G_SUCCESS); - }); - - return {context, lifetime}; +void GPUStreamOwner::destroy() noexcept { + if (m_stream != nullptr) { + g_stream_destroy(m_stream); + m_stream = nullptr; + } } -GPUContextHandle GPUContextHandle::retain_primary_context_for_device(g_device_t device) { - g_context_t context; - KMM_GPU_CHECK(g_device_primary_ctx_retain(&context, device)); +std::ostream& operator<<(std::ostream& stream, const GPUStreamOwner& self) { + if (self.m_stream == nullptr) { + return stream << "GPU-stream: none"; + } - auto lifetime = std::shared_ptr(nullptr, [=](const void* ignore) { - KMM_ASSERT(g_device_primary_ctx_release(device) == G_SUCCESS); - }); + return stream << GPUStreamRef(self.m_stream); +} - return {context, lifetime}; +std::ostream& operator<<(std::ostream& stream, const GPUStreamRef& self) { + return stream << self.m_stream_id; } -GPUContextGuard::GPUContextGuard(GPUContextHandle context) : m_context(std::move(context)) { - KMM_GPU_CHECK(g_ctx_push_current(m_context)); +std::ostream& operator<<(std::ostream& stream, const GPUStreamId& self) { + return stream << "GPU-stream:" << self.m_id; } -GPUContextGuard::~GPUContextGuard() { - g_context_t previous; - KMM_ASSERT(g_ctx_pop_current(&previous) == G_SUCCESS); +std::ostream& operator<<(std::ostream& stream, const GPUContextId& self) { + return stream << "GPU-context:" << self.m_id; } } // namespace kmm diff --git a/src/utils/notify.cpp b/src/utils/notify.cpp index 3b41fcf8..1ca2d021 100644 --- a/src/utils/notify.cpp +++ b/src/utils/notify.cpp @@ -2,13 +2,7 @@ namespace kmm { -NotifyHandle::NotifyHandle(std::shared_ptr m) : m_impl(std::move(m)) {} - -NotifyHandle::NotifyHandle(std::unique_ptr m) : m_impl(std::move(m)) {} - -NotifyHandle::~NotifyHandle() { - notify_and_clear(); -} +NotifyHandle::~NotifyHandle() = default; void NotifyHandle::notify() const noexcept { if (m_impl) { @@ -17,13 +11,12 @@ void NotifyHandle::notify() const noexcept { } void NotifyHandle::clear() noexcept { - m_impl = nullptr; + m_impl.reset(); } void NotifyHandle::notify_and_clear() noexcept { - if (auto m = std::exchange(m_impl, nullptr)) { - m->notify(); - } + notify(); + clear(); } -} // namespace kmm \ No newline at end of file +} // namespace kmm diff --git a/src/utils/small_vector.cpp b/src/utils/small_vector.cpp index 92074151..5f40685c 100644 --- a/src/utils/small_vector.cpp +++ b/src/utils/small_vector.cpp @@ -1,9 +1,7 @@ -#include - #include "kmm/utils/small_vector.hpp" namespace kmm { [[noreturn]] __attribute__((noinline)) void throw_small_vector_out_of_capacity() { - throw std::overflow_error("small_vector exceeds capacity"); + throw std::runtime_error("small_vector exceeds capacity"); } } // namespace kmm \ No newline at end of file diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index fee9e8e9..193cfe4b 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -1,6 +1,6 @@ file(GLOB_RECURSE sources - "${PROJECT_SOURCE_DIR}/test/*.cpp" - "${PROJECT_SOURCE_DIR}/test/*.cu") + "${PROJECT_SOURCE_DIR}/test/*.cpp" + "${PROJECT_SOURCE_DIR}/test/*.cu") add_executable(kmmTest ${sources}) @@ -12,4 +12,7 @@ target_compile_features(kmmTest PRIVATE cxx_std_17) target_link_libraries(kmmTest PRIVATE kmm) target_link_libraries(kmmTest PRIVATE Catch2::Catch2WithMain) -include(Catch) \ No newline at end of file +find_package(Threads REQUIRED) +target_link_libraries(kmmTest PRIVATE Threads::Threads) + +include(Catch) diff --git a/test/api/test_context_sum.cpp b/test/api/test_context_sum.cpp new file mode 100644 index 00000000..689609de --- /dev/null +++ b/test/api/test_context_sum.cpp @@ -0,0 +1,114 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memops/reduction.hpp" + +using namespace kmm; +using namespace kmm::memops; + +namespace { +// Minimal duck-typed stand-in for `kmm::Layout`, avoiding a dependency on `kmm/core/layout.hpp` +// (mirrors `FakeLayout` in test/memops/test_copy.cpp). +template +struct FakeLayout { + static constexpr size_t rank = N; + + ptrdiff_t offset = 0; + ptrdiff_t extents[N]; + ptrdiff_t strides[N]; + ptrdiff_t origins[N] = {}; + + ptrdiff_t base_offset() const { + return offset; + } + + ptrdiff_t extent(size_t axis) const { + return extents[axis]; + } + + ptrdiff_t stride(size_t axis) const { + return strides[axis]; + } + + ptrdiff_t begin(size_t axis) const { + return origins[axis]; + } +}; +} // namespace + +TEST_CASE("make_reduction_description") { + // A contiguous row-major (4, 3) source reduced over its trailing axis into a contiguous + // 4-element destination. + FakeLayout<1> dst {/* offset */ 0, /* extents */ {4}, /* strides */ {1}}; + FakeLayout<2> src {/* offset */ 0, /* extents */ {4, 3}, /* strides */ {3, 1}}; + + ReductionDescription description = + make_reduction_description(dst, src, 1, DataType::Float32, ReductionOp::Sum); + + CHECK(description.dtype == DataType::Float32); + CHECK(description.operation == ReductionOp::Sum); + CHECK(description.input_offset == 0); + CHECK(description.output_offset == 0); + CHECK(description.reduction_extent == 3); + CHECK(description.reduction_stride == static_cast(sizeof(float))); + + REQUIRE(description.num_dims == 1); + CHECK(description.dims[0].extent == 4); + CHECK(description.dims[0].input_stride == 3 * static_cast(sizeof(float))); + CHECK(description.dims[0].output_stride == static_cast(sizeof(float))); +} + +TEST_CASE("make_reduction_description (leading axis)") { + // The same (4, 3) source, but now reduced over its leading axis into a contiguous + // 3-element destination. + FakeLayout<1> dst {/* offset */ 0, /* extents */ {3}, /* strides */ {1}}; + FakeLayout<2> src {/* offset */ 0, /* extents */ {4, 3}, /* strides */ {3, 1}}; + + ReductionDescription description = + make_reduction_description(dst, src, 0, DataType::Float32, ReductionOp::Sum); + + CHECK(description.reduction_extent == 4); + CHECK(description.reduction_stride == 3 * static_cast(sizeof(float))); + + REQUIRE(description.num_dims == 1); + CHECK(description.dims[0].extent == 3); + CHECK(description.dims[0].input_stride == static_cast(sizeof(float))); + CHECK(description.dims[0].output_stride == static_cast(sizeof(float))); +} + +TEST_CASE("make_reduction_description end-to-end (CPU)") { + // Reduce a (4, 3) row-major input down to 4 outputs by summing over the trailing axis, going + // through `make_reduction_description` instead of building the description by hand (compare + // to the "reduce (CPU)" case in test/memops/test_reduction.cpp). + std::vector src = { + 1, + 2, + 3, // + 4, + 5, + 6, // + 7, + 8, + 9, // + 10, + 11, + 12 // + }; + std::vector dst(4, 0.0f); + + FakeLayout<1> dst_layout {/* offset */ 0, /* extents */ {4}, /* strides */ {1}}; + FakeLayout<2> src_layout {/* offset */ 0, /* extents */ {4, 3}, /* strides */ {3, 1}}; + + ReductionDescription description = make_reduction_description( // + dst_layout, + src_layout, + 1, + DataType::Float32, + ReductionOp::Sum + ); + + reduce(src.data(), dst.data(), description); + + CHECK(dst == std::vector {6, 15, 24, 33}); +} diff --git a/test/core/test_bounds.cpp b/test/core/test_bounds.cpp new file mode 100644 index 00000000..775ef098 --- /dev/null +++ b/test/core/test_bounds.cpp @@ -0,0 +1,179 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/bounds.hpp" + +using namespace kmm; + +TEST_CASE("Bounds construction") { + SECTION("default constructs all-empty ranges") { + Bounds<2, int> b; + CHECK(b[0] == Range()); + CHECK(b[1] == Range()); + CHECK(b.is_empty()); + } + + SECTION("variadic constructor takes exactly N ranges") { + Bounds<2, int> b(Range(0, 3), Range(1, 4)); + CHECK(b[0] == Range(0, 3)); + CHECK(b[1] == Range(1, 4)); + } + + SECTION("bounds() free function") { + auto b = bounds(Range(0, 3), Range(1, 4)); + CHECK(b[0] == Range(0, 3)); + CHECK(b[1] == Range(1, 4)); + } + + SECTION("wraps an existing Vec, N> unchanged") { + Vec, 2> v {Range(0, 3), Range(1, 4)}; + Bounds<2, int> b(v); + CHECK(b[0] == Range(0, 3)); + CHECK(b[1] == Range(1, 4)); + } + + SECTION("constructs from a Shape as the region [0, shape)") { + Bounds<2, int> b(Shape<2, int>(3, 4)); + CHECK(b.begin() == Point<2, int>(0, 0)); + CHECK(b.end() == Point<2, int>(3, 4)); + } +} + +TEST_CASE("Bounds::from_bounds") { + Bounds<2, int> b = Bounds<2, int>::from_bounds(Point<2, int>(1, 2), Point<2, int>(4, 6)); + + CHECK(b.begin() == Point<2, int>(1, 2)); + CHECK(b.end() == Point<2, int>(4, 6)); +} + +TEST_CASE("Bounds::from_offset_size") { + Bounds<2, int> b = Bounds<2, int>::from_offset_size(Point<2, int>(1, 2), Shape<2, int>(3, 4)); + + CHECK(b.begin() == Point<2, int>(1, 2)); + CHECK(b.end() == Point<2, int>(4, 6)); + CHECK(b.shape() == Shape<2, int>(3, 4)); +} + +TEST_CASE("Bounds::empty/one") { + CHECK(Bounds<2, int>::empty().is_empty()); + CHECK_FALSE(Bounds<2, int>::one().is_empty()); + CHECK(Bounds<2, int>::one().shape() == Shape<2, int>(1, 1)); +} + +TEST_CASE("Bounds::from") { + SECTION("same dimensionality copies ranges") { + Vec, 2> v {Range(0, 3), Range(1, 4)}; + Bounds<2, int> b = Bounds<2, int>::from(v); + CHECK(b[0] == Range(0, 3)); + CHECK(b[1] == Range(1, 4)); + } + + SECTION("growing dimensionality pads with the unit range 0...1") { + Vec, 1> v {Range(0, 3)}; + Bounds<2, int> b = Bounds<2, int>::from(v); + CHECK(b[0] == Range(0, 3)); + CHECK(b[1] == Range(0, 1)); + } +} + +TEST_CASE("Bounds converting constructor") { + SECTION("growing dimensionality pads with the unit range") { + Bounds<1, int> src(Range(0, 3)); + Bounds<2, int> dst(src); + CHECK(dst[0] == Range(0, 3)); + CHECK(dst[1] == Range(0, 1)); + } + + SECTION("shrinking dimensionality succeeds when dropped axes are the unit range") { + Bounds<2, int> src(Range(0, 3), Range(0, 1)); + Bounds<1, int> dst(src); + CHECK(dst[0] == Range(0, 3)); + } + + SECTION("shrinking dimensionality throws when a dropped axis is not the unit range") { + Bounds<2, int> src(Range(0, 3), Range(1, 4)); + CHECK_THROWS_AS((Bounds<1, int>(src)), std::overflow_error); + } +} + +TEST_CASE("Bounds::begin/end/size accessors") { + Bounds<2, int> b(Range(1, 4), Range(2, 6)); + + CHECK(b.begin(0) == 1); + CHECK(b.end(0) == 4); + CHECK(b.size(0) == 3); + + CHECK(b.begin(1) == 2); + CHECK(b.end(1) == 6); + CHECK(b.size(1) == 4); + + SECTION("out-of-range axis defaults to the unit range") { + CHECK(b.begin(2) == 0); + CHECK(b.end(2) == 1); + CHECK(b.size(2) == 1); + } +} + +TEST_CASE("Bounds::is_empty") { + CHECK_FALSE(Bounds<2, int>(Range(0, 3), Range(0, 4)).is_empty()); + CHECK(Bounds<2, int>(Range(0, 0), Range(0, 4)).is_empty()); + CHECK(Bounds<2, int>(Range(3, 0), Range(0, 4)).is_empty()); +} + +TEST_CASE("Bounds::volume") { + CHECK(Bounds<2, int>(Range(0, 3), Range(0, 4)).volume() == 12); + CHECK(Bounds<2, int>(Range(0, 0), Range(0, 4)).volume() == 0); +} + +TEST_CASE("Bounds::intersection") { + Bounds<1, int> a(Range(0, 10)); + Bounds<1, int> b(Range(5, 15)); + + CHECK(a.intersection(b)[0] == Range(5, 10)); + + Bounds<1, int> c(Range(20, 30)); + CHECK(a.intersection(c).is_empty()); +} + +TEST_CASE("Bounds::overlaps") { + Bounds<1, int> a(Range(0, 10)); + + CHECK(a.overlaps(Bounds<1, int>(Range(5, 15)))); + CHECK_FALSE(a.overlaps(Bounds<1, int>(Range(10, 20)))); + CHECK_FALSE(a.overlaps(Bounds<1, int>(Range(5, 5)))); + + SECTION("overlaps(Shape)") { + CHECK(a.overlaps(Shape<1, int>(5))); + CHECK_FALSE(a.overlaps(Shape<1, int>(0))); + } +} + +TEST_CASE("Bounds::contains(Bounds)") { + Bounds<1, int> outer(Range(0, 10)); + + CHECK(outer.contains(Bounds<1, int>(Range(2, 8)))); + CHECK(outer.contains(Bounds<1, int>(Range(0, 10)))); + CHECK_FALSE(outer.contains(Bounds<1, int>(Range(-1, 8)))); + CHECK_FALSE(outer.contains(Bounds<1, int>(Range(2, 11)))); + CHECK(outer.contains(Bounds<1, int>(Range(5, 5)))); // empty range always contained + + SECTION("contains(Shape)") { + CHECK(outer.contains(Shape<1, int>(10))); + CHECK_FALSE(outer.contains(Shape<1, int>(11))); + } +} + +TEST_CASE("Bounds::contains(Point)") { + Bounds<2, int> b(Range(0, 3), Range(0, 4)); + + CHECK(b.contains(Point<2, int>(0, 0))); + CHECK(b.contains(Point<2, int>(2, 3))); + CHECK_FALSE(b.contains(Point<2, int>(3, 0))); + CHECK_FALSE(b.contains(Point<2, int>(0, 4))); + + SECTION("variadic overload") { + CHECK(b.contains(1, 1)); + CHECK_FALSE(b.contains(3, 3)); + } +} diff --git a/test/core/test_checked_compare.cpp b/test/core/test_checked_compare.cpp new file mode 100644 index 00000000..c5f5cca0 --- /dev/null +++ b/test/core/test_checked_compare.cpp @@ -0,0 +1,338 @@ +#include +#include +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/checked_compare.hpp" + +using namespace kmm; + +using u8 = uint8_t; +using u16 = uint16_t; +using u32 = uint32_t; +using u64 = uint64_t; + +using i8 = int8_t; +using i16 = int16_t; +using i32 = int32_t; +using i64 = int64_t; + +using f32 = float; +using f64 = double; + +#define CHECK_IS_LESS(A, B) \ + CHECK(is_less(A, B)); \ + CHECK(is_less_equal(A, B)); \ + CHECK_FALSE(is_greater_equal(A, B)); \ + CHECK_FALSE(is_greater(A, B)); \ + CHECK_FALSE(is_less(B, A)); \ + CHECK_FALSE(is_less_equal(B, A)); \ + CHECK(is_greater(B, A)); \ + CHECK(is_greater_equal(B, A)); + +TEST_CASE("is_less") { + SECTION("signed vs unsigned, negative") { + CHECK_IS_LESS(i8(-128), u8(0)); + CHECK_IS_LESS(i16(-1), u16(1)); + CHECK_IS_LESS(i32(-100), u32(0)); + CHECK_IS_LESS(i64(std::numeric_limits::min()), u64(0ULL)); + CHECK_IS_LESS(i32(-1), u32(10)); + CHECK_IS_LESS(i32(-1), i32(0)); + } + + SECTION("signed vs unsigned, positive") { + CHECK_IS_LESS(u8(0), i8(127)); + CHECK_IS_LESS(u16(1), i16(2)); + CHECK_IS_LESS(u32(0), i64(1)); + CHECK_IS_LESS(u64(0ULL), i64(LONG_MAX)); + CHECK_IS_LESS(i32(INT_MIN), i32(0)); + CHECK_IS_LESS(i32(-2147483647), u32(2147483648u)); + CHECK_IS_LESS(i64(0), u64(ULONG_MAX)); + } + + SECTION("different integers") { + CHECK_IS_LESS(u8(200), u16(300)); + CHECK_IS_LESS(i16(30000), i32(40000)); + CHECK_IS_LESS(i32(2147483646), i64(2147483647)); + CHECK_IS_LESS(u32(4294967294U), u64(4294967295ULL)); + CHECK_IS_LESS(i64(9223372036854775806LL), u64(9223372036854775807ULL)); + } + + SECTION("float vs float") { + CHECK_IS_LESS(f32(1.5f), f64(1.6)); + CHECK_IS_LESS(f64(-1.0), f32(0.0f)); + CHECK_IS_LESS(f32(-std::numeric_limits::max()), f64(std::numeric_limits::max())); + CHECK_IS_LESS(f64(-std::numeric_limits::infinity()), f64(-1e308)); + CHECK_IS_LESS(f32(-1e10f), f64(-1e9)); + CHECK_IS_LESS(f64(1.1), f64(1.2L)); + CHECK_IS_LESS(f64(-1.2L), f64(-1.1)); + CHECK_IS_LESS(f64(0.0L), f64(0.1L)); + CHECK_IS_LESS(f64(0.0), std::numeric_limits::infinity()); + CHECK_IS_LESS(f32(-std::numeric_limits::infinity()), f32(0.0f)); + CHECK_IS_LESS(f32(std::numeric_limits::max()), f64(std::numeric_limits::max())); + CHECK_IS_LESS(f32(1.0001f), f64(1.0002)); + } + + SECTION("float vs integer, fractional values") { + CHECK_IS_LESS(u16(0), f64(0.1)); + CHECK_IS_LESS(u64(100), f64(101.0)); + CHECK_IS_LESS(i64(-1000), f32(0.0f)); + CHECK_IS_LESS(f32(0.999f), i32(1)); + CHECK_IS_LESS(i8(-10), f32(-9.9f)); + CHECK_IS_LESS(i32(1), f64(1.1)); + CHECK_IS_LESS(i64(-1), f64(0.0)); + CHECK_IS_LESS(u32(1), f64(1.1)); + CHECK_IS_LESS(f64(-0.01), u32(0)); + CHECK_IS_LESS(f32(-0.1f), u16(0)); + CHECK_IS_LESS(u64(1), f64(1.0001L)); + CHECK_IS_LESS(i8(10), f64(10.1L)); + CHECK_IS_LESS(i16(123), f64(123.1L)); + CHECK_IS_LESS(u16(65535), f64(65535.5)); + CHECK_IS_LESS(i32(-1), f64(0.0L)); + + // f64(LONG_MAX) will round up, f64(LONG_MIN + 1) will round down. + CHECK_IS_LESS(LONG_MAX, f64(LONG_MAX)); + CHECK_IS_LESS(f64(LONG_MIN + 1), LONG_MIN + 1); + } + + SECTION("float rules: equal values, -0.0 vs 0.0, NaN") { + CHECK_FALSE(is_less(-0.0, 0.0)); + CHECK_FALSE(is_less(f32(1.1f), f64(1.1))); + CHECK_FALSE(is_less(u8(255), u8(255))); + CHECK_FALSE(is_less(std::numeric_limits::quiet_NaN(), f64(0.0))); + } +} + +#define CHECK_IS_EQUAL(A, B) \ + CHECK(is_equal(A, B)); \ + CHECK(is_equal(B, A)); \ + CHECK_FALSE(is_less(A, B)); \ + CHECK_FALSE(is_less(B, A)); \ + CHECK(is_less_equal(A, B)); \ + CHECK(is_less_equal(B, A)); \ + CHECK(is_greater_equal(A, B)); \ + CHECK(is_greater_equal(B, A)); \ + CHECK_FALSE(is_greater(A, B)); \ + CHECK_FALSE(is_greater(B, A)); + +#define CHECK_NOT_EQUAL(A, B) \ + CHECK_FALSE(is_equal(A, B)); \ + CHECK_FALSE(is_equal(B, A)); \ + CHECK((is_less(A, B) || is_less(B, A))); + +TEST_CASE("is_equal") { + SECTION("signed vs unsigned") { + CHECK_IS_EQUAL(i8(0), u8(0)); + CHECK_IS_EQUAL(i16(123), u16(123)); + CHECK_IS_EQUAL(u32(3000), i64(3000)); + CHECK_IS_EQUAL(u64(0ULL), i64(0LL)); + CHECK_IS_EQUAL(u8(255), i32(255)); + CHECK_IS_EQUAL(u32(1), f64(1.0)); + CHECK_IS_EQUAL((unsigned long long)(42), f64(42.0L)); + CHECK_IS_EQUAL(i64(1234567890123LL), u64(1234567890123ULL)); + CHECK_IS_EQUAL(u32(0), f32(0.0f)); + } + + SECTION("different integers") { + CHECK_IS_EQUAL(i32(-100), i64(-100)); + CHECK_IS_EQUAL(i32(INT_MAX), i64(INT_MAX)); + CHECK_IS_EQUAL(i8(-128), i16(-128)); + CHECK_IS_EQUAL(u16(65535), u32(65535)); + CHECK_IS_EQUAL(i16(-32768), i32(-32768)); + } + + SECTION("float vs float") { + CHECK_IS_EQUAL(f32(2.5f), f64(2.5)); + CHECK_IS_EQUAL(f64(-1.0), f64(-1.0L)); + CHECK_IS_EQUAL( + f32(std::numeric_limits::infinity()), + f64(std::numeric_limits::infinity()) + ); + CHECK_IS_EQUAL(f64(-0.0), f32(-0.0f)); + CHECK_IS_EQUAL(f64(0.0L), f64(0.0)); + CHECK_IS_EQUAL(f32(0.5f), f64(0.5L)); + CHECK_IS_EQUAL( + f64(-std::numeric_limits::infinity()), + f64(-std::numeric_limits::infinity()) + ); + CHECK_IS_EQUAL(f64(123456.0L), f64(123456.0L)); + } + + SECTION("float vs integer") { + CHECK_IS_EQUAL(i64(-1), f64(-1.0L)); + CHECK_IS_EQUAL(i32(2147483647), f64(2147483647.0)); + } + + SECTION("signed vs unsigned, not equal") { + CHECK_NOT_EQUAL(i8(0), u8(1)); + CHECK_NOT_EQUAL(i16(123), u16(124)); + CHECK_NOT_EQUAL(i32(-100), i64(100)); + CHECK_NOT_EQUAL(u32(3000), i64(-3000)); + CHECK_NOT_EQUAL(u64(1), i64(2)); + CHECK_NOT_EQUAL(u8(255), i32(254)); + CHECK_NOT_EQUAL(u32(1), f64(2.0)); + CHECK_NOT_EQUAL((unsigned long long)(42), f64(42.1)); + } + + SECTION("integers, not equal") { + CHECK_NOT_EQUAL(i32(INT_MAX), i64(INT_MAX) - 1); + CHECK_NOT_EQUAL(i8(-128), i16(-127)); + CHECK_NOT_EQUAL(u16(65534), u32(65535)); + CHECK_NOT_EQUAL(i16(-32768), i32(-32767)); + CHECK_NOT_EQUAL(i64(1234567890123LL), u64(1234567890124ULL)); + } + + SECTION("floats, not equal") { + CHECK_NOT_EQUAL(f32(2.5f), f64(2.5001)); + CHECK_NOT_EQUAL(f64(-1.0), f64(-1.0000000001L)); + CHECK_NOT_EQUAL(u64(ULONG_MAX), f64(ULONG_MAX) - 1.0); + CHECK_NOT_EQUAL(u64(ULONG_MAX), f64(ULONG_MAX)); + CHECK_NOT_EQUAL( + f32(std::numeric_limits::infinity()), + f64(std::numeric_limits::lowest()) + ); + CHECK_NOT_EQUAL(f64(123.0), f64(124.0)); + CHECK_NOT_EQUAL(f64(1.0L), f64(1.0000000001L)); + CHECK_NOT_EQUAL(f64(0.0L), f64(0.1)); + CHECK_NOT_EQUAL(f32(0.5f), f64(0.5000001L)); + CHECK_NOT_EQUAL(u32(0), f32(0.0001f)); + } + + SECTION("float vs integer, not equal") { + CHECK_NOT_EQUAL(i64(-1), f64(1.0L)); + CHECK_NOT_EQUAL(i32(2147483647), f64(2147483646.0)); + } + + SECTION("NaN is never equal") { + CHECK_FALSE(is_equal(std::numeric_limits::quiet_NaN(), f64(0.0))); + CHECK_FALSE(is_equal(std::numeric_limits::quiet_NaN(), i32(0))); + CHECK_FALSE( + is_equal(std::numeric_limits::quiet_NaN(), std::numeric_limits::quiet_NaN()) + ); + } +} + +TEST_CASE("is_convertible") { + SECTION("convertible: integer to integer") { + CHECK(is_convertible(i32(0))); + CHECK(is_convertible(i64(4294967295LL))); + CHECK(is_convertible(i8(127))); + CHECK(is_convertible(i32(255))); + CHECK(is_convertible(i16(-128))); + CHECK(is_convertible(u32(65535))); + CHECK(is_convertible(i32(INT_MIN))); + CHECK(is_convertible(u32(100))); + CHECK(is_convertible(u64(12345ULL))); + CHECK(is_convertible(i64(0))); + } + + SECTION("convertible: integer to float") { + CHECK(is_convertible(i32(123456))); + CHECK(is_convertible(f32(1.5f))); + CHECK(is_convertible(f64(1.5))); + CHECK(is_convertible(INT_MIN)); + CHECK(is_convertible(std::numeric_limits::infinity())); + CHECK(is_convertible(1.0 + std::numeric_limits::epsilon())); + } + + SECTION("convertible: float to integer, whole value") { + CHECK(is_convertible(f64(0.0))); + CHECK(is_convertible(f64(2147483647.0))); + CHECK(is_convertible(f64(9007199254740992.0))); + CHECK(is_convertible(f64(0.0))); + CHECK(is_convertible(f32(-100.0f))); + CHECK(is_convertible(f64(1.0L))); + CHECK(is_convertible(f64(-2147483648.0))); + CHECK(is_convertible(f64(9007199254740991.0))); + CHECK(is_convertible(f32(16777216.0f))); + CHECK(is_convertible(f64(0.0))); + CHECK(is_convertible(f64(-32768.0))); + } + + SECTION("not convertible: integer overflow/underflow") { + CHECK_FALSE(is_convertible(i32(-1))); + CHECK_FALSE(is_convertible(i16(128))); + CHECK_FALSE(is_convertible(i32(70000))); + CHECK_FALSE(is_convertible(i32(256))); + CHECK_FALSE(is_convertible(i32(-1))); + } + + SECTION("not convertible: fractional value") { + CHECK_FALSE(is_convertible(f64(0.1))); + CHECK_FALSE(is_convertible(f64(2147483648.0))); + CHECK_FALSE(is_convertible(f64(255.5))); + CHECK_FALSE(is_convertible(f64(40000.0))); + CHECK_FALSE(is_convertible(f64(9.223372036854775808e18))); + CHECK_FALSE(is_convertible(f64(-0.1))); + CHECK_FALSE(is_convertible(f32(2147483648.0f))); + CHECK_FALSE(is_convertible(f64(-32769.0))); + CHECK_FALSE(is_convertible(f64(-1.0))); + CHECK_FALSE(is_convertible(f64(1.5L))); + CHECK_FALSE(is_convertible(f64(std::numeric_limits::max()))); + CHECK_FALSE(is_convertible(f64(-2147483649.0))); + CHECK_FALSE(is_convertible(f32(1.0000001f))); + CHECK_FALSE(is_convertible(f64(-0.00001))); + CHECK_FALSE(is_convertible(f64(1e20))); + CHECK_FALSE(is_convertible(f32(32768.0f))); + CHECK_FALSE(is_convertible(1.0 + std::numeric_limits::epsilon())); + CHECK_FALSE(is_convertible(1.0f + std::numeric_limits::epsilon())); + CHECK_FALSE(is_convertible(std::numeric_limits::epsilon())); + } + + SECTION("not convertible: infinity and NaN") { + CHECK_FALSE(is_convertible(f64(std::numeric_limits::infinity()))); + CHECK_FALSE(is_convertible(f64(std::numeric_limits::quiet_NaN()))); + CHECK_FALSE(is_convertible(f64(std::numeric_limits::infinity()))); + CHECK_FALSE(is_convertible(f64(std::numeric_limits::quiet_NaN()))); + } +} + +TEST_CASE("checked_cast") { + SECTION("widening nothrow") { + CHECK_NOTHROW(checked_cast(INT_MAX)); + CHECK_NOTHROW(checked_cast(UINT_MAX)); + CHECK_NOTHROW(checked_cast(LONG_MAX)); + CHECK_NOTHROW(checked_cast(ULONG_MAX)); + + CHECK_NOTHROW(checked_cast(INT_MAX)); + CHECK_NOTHROW(checked_cast(UINT_MAX)); + CHECK_NOTHROW(checked_cast(LONG_MAX)); + + CHECK_NOTHROW(checked_cast(INT_MAX)); + CHECK_NOTHROW(checked_cast(UINT_MAX)); + + CHECK_NOTHROW(checked_cast(INT_MAX)); + } + + SECTION("narrowing throw") { + CHECK_THROWS(checked_cast(ULONG_MAX)); + + CHECK_THROWS(checked_cast(LONG_MAX)); + CHECK_THROWS(checked_cast(ULONG_MAX)); + + CHECK_THROWS(checked_cast(UINT_MAX)); + CHECK_THROWS(checked_cast(LONG_MAX)); + CHECK_THROWS(checked_cast(ULONG_MAX)); + + CHECK_THROWS(checked_cast(i32(200))); + CHECK_THROWS(checked_cast(f64(200.0))); + CHECK_THROWS(checked_cast(i64(-1337))); + } + + SECTION("float to integer nothrow") { + CHECK_NOTHROW(checked_cast(i32(200))); + CHECK_NOTHROW(checked_cast(f64(201.0))); + CHECK_NOTHROW(checked_cast(i32(202))); + } + + SECTION("integer to float throw") { + CHECK_THROWS(checked_cast(INT_MAX)); // not representable + CHECK_THROWS(checked_cast(5.5)); // fractional part + CHECK_THROWS(checked_cast(-5.0)); // negative value + CHECK_THROWS(checked_cast(std::numeric_limits::max())); + CHECK_THROWS(checked_cast(std::numeric_limits::infinity())); + CHECK_THROWS(checked_cast(std::numeric_limits::quiet_NaN())); + } +} \ No newline at end of file diff --git a/test/core/test_checked_math.cpp b/test/core/test_checked_math.cpp new file mode 100644 index 00000000..6a37616c --- /dev/null +++ b/test/core/test_checked_math.cpp @@ -0,0 +1,361 @@ +#include +#include +#include +#include +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/checked_math.hpp" + +using namespace kmm; + +TEST_CASE("checked_add") { + CHECK(checked_add(2U, 3U) == 5); + CHECK(checked_add(-2, 3) == 1); + CHECK(checked_add(2UL, 3UL) == 5); + CHECK(checked_add(-2L, 3L) == 1); + + CHECK(checked_add(INT_MAX - 1, 1) == INT_MAX); + CHECK(checked_add(INT_MIN + 1, -1) == INT_MIN); + CHECK(checked_add(UINT_MAX - 1, 1U) == UINT_MAX); + + CHECK_THROWS(checked_add(INT_MAX, 1)); + CHECK_THROWS(checked_add(INT_MIN, -1)); + CHECK_THROWS(checked_add(UINT_MAX, 1U)); + + CHECK_THROWS(checked_add(INT_MAX, INT_MAX)); + CHECK_THROWS(checked_add(INT_MIN, INT_MIN)); + CHECK_THROWS(checked_add(UINT_MAX, UINT_MAX)); + + CHECK(checked_add(LONG_MAX - 1, 1L) == LONG_MAX); + CHECK(checked_add(LONG_MIN + 1, -1L) == LONG_MIN); + CHECK(checked_add(ULONG_MAX - 1, 1LU) == ULONG_MAX); + + CHECK_THROWS(checked_add(LONG_MAX, 1L)); + CHECK_THROWS(checked_add(LONG_MIN, -1L)); + CHECK_THROWS(checked_add(ULONG_MAX, 1UL)); + + CHECK_THROWS(checked_add(LONG_MAX, LONG_MAX)); + CHECK_THROWS(checked_add(LONG_MIN, LONG_MIN)); + CHECK_THROWS(checked_add(ULONG_MAX, ULONG_MAX)); +} + +TEST_CASE("checked_sub") { + CHECK(checked_sub(5U, 3U) == 2); + CHECK(checked_sub(3, 5) == -2); + CHECK(checked_sub(5UL, 3UL) == 2); + CHECK(checked_sub(3L, 5L) == -2L); + + CHECK(checked_sub(INT_MAX - 1, -1) == INT_MAX); + CHECK(checked_sub(INT_MIN + 1, 1) == INT_MIN); + CHECK(checked_sub(UINT_MAX, UINT_MAX) == 0U); + + CHECK_THROWS(checked_sub(INT_MAX, -1)); + CHECK_THROWS(checked_sub(INT_MIN, 1)); + CHECK_THROWS(checked_sub(0U, 1U)); + + CHECK_THROWS(checked_sub(INT_MAX, INT_MIN)); + CHECK_THROWS(checked_sub(INT_MIN, INT_MAX)); + CHECK_THROWS(checked_sub(0U, UINT_MAX)); + + CHECK(checked_sub(LONG_MAX - 1L, -1L) == LONG_MAX); + CHECK(checked_sub(LONG_MIN + 1L, 1L) == LONG_MIN); + CHECK(checked_sub(ULONG_MAX, ULONG_MAX) == 0UL); + + CHECK_THROWS(checked_sub(LONG_MAX, -1L)); + CHECK_THROWS(checked_sub(LONG_MIN, 1L)); + CHECK_THROWS(checked_sub(0UL, 1UL)); + + CHECK_THROWS(checked_sub(LONG_MAX, LONG_MIN)); + CHECK_THROWS(checked_sub(LONG_MIN, LONG_MAX)); + CHECK_THROWS(checked_sub(0UL, ULONG_MAX)); +} + +TEST_CASE("checked_mul") { + CHECK(checked_mul(10U, 3U) == 30); + CHECK(checked_mul(-10, 3) == -30); + CHECK(checked_mul(10UL, 3UL) == 30); + CHECK(checked_mul(-10L, 3L) == -30L); + CHECK(checked_mul(10, 0) == 0); + CHECK(checked_mul(-10, 0) == 0); + + CHECK(checked_mul(INT_MAX, 1) == INT_MAX); + CHECK(checked_mul(INT_MAX, -1) == INT_MIN + 1); + CHECK(checked_mul(INT_MIN, 1) == INT_MIN); + CHECK(checked_mul(UINT_MAX, 1U) == UINT_MAX); + + CHECK_THROWS(checked_mul(INT_MAX, 2)); + CHECK_THROWS(checked_mul(INT_MIN, 2)); + CHECK_THROWS(checked_mul(INT_MIN, -1)); + CHECK_THROWS(checked_mul(UINT_MAX, 2U)); + + CHECK_THROWS(checked_mul(INT_MAX, INT_MAX)); + CHECK_THROWS(checked_mul(INT_MIN, INT_MIN)); + CHECK_THROWS(checked_mul(UINT_MAX, UINT_MAX)); + + CHECK(checked_mul(LONG_MAX, 1L) == LONG_MAX); + CHECK(checked_mul(LONG_MIN, 1L) == LONG_MIN); + CHECK(checked_mul(ULONG_MAX, 1UL) == ULONG_MAX); + + CHECK_THROWS(checked_mul(LONG_MAX, 2L)); + CHECK_THROWS(checked_mul(LONG_MIN, 2L)); + CHECK_THROWS(checked_mul(ULONG_MAX, 2UL)); + + CHECK_THROWS(checked_mul(LONG_MAX, LONG_MIN)); + CHECK_THROWS(checked_mul(LONG_MIN, LONG_MAX)); + CHECK_THROWS(checked_mul(ULONG_MAX, ULONG_MAX)); +} + +TEST_CASE("checked_div") { + CHECK(checked_div(10, 3) == 3); + CHECK(checked_div(-10, 3) == -3); + + // division by zero + CHECK_THROWS(checked_div(10, 0)); + + // overflow + CHECK_THROWS(checked_div(INT_MIN, -1)); +} + +TEST_CASE("checked_rem") { + CHECK(checked_rem(10, 3) == 1); + CHECK(checked_rem(-10, 3) == -1); + CHECK_THROWS(checked_rem(10, 0)); +} + +TEST_CASE("checked_neg") { + CHECK(checked_neg(5) == -5); + CHECK(checked_neg(-5) == 5); + CHECK_THROWS(checked_neg(INT_MIN)); +} + +TEST_CASE("checked_abs") { + CHECK(checked_abs(5) == 5); + CHECK(checked_abs(-5) == 5); + CHECK(checked_abs(0) == 0); + CHECK_THROWS(checked_abs(INT_MIN)); +} + +TEST_CASE("checked_sum") { + std::vector a = {}; + CHECK(checked_sum(a.begin(), a.end()) == 0); + + a = {1, 2, 3, 4}; + CHECK(checked_sum(a.begin(), a.end()) == 10); + + a = {1, 2, 3, 4}; + CHECK(checked_sum(a.begin(), a.end(), 5) == 15); + + a = {1, 2, 3, 4, INT_MAX}; + CHECK_THROWS(checked_sum(a.begin(), a.end())); + + a = {1, 2, 3, 4, INT_MAX}; + CHECK(checked_sum(a.begin(), a.end(), long()) == long(INT_MAX) + 10); +} + +TEST_CASE("checked_product") { + std::vector a = {}; + CHECK(checked_product(a.begin(), a.end()) == 1); + + a = {1, 2, 3, 4}; + CHECK(checked_product(a.begin(), a.end()) == 24); + + a = {1, 2, 3, 4}; + CHECK(checked_product(a.begin(), a.end(), 5) == 120); + + a = {1, 2, 3, 4, INT_MAX}; + CHECK_THROWS(checked_product(a.begin(), a.end())); + + a = {1, 2, 3, 4, INT_MAX}; + CHECK(checked_product(a.begin(), a.end(), long(1)) == long(INT_MAX) * 24); +} + +template +std::vector generate_inputs(std::mt19937_64 rng, size_t n) { + std::vector inputs = { + std::numeric_limits::min(), + std::numeric_limits::min() + static_cast(1), + std::numeric_limits::min() + static_cast(2), + std::numeric_limits::min() + static_cast(3), + std::numeric_limits::min() / static_cast(2) + static_cast(1), + std::numeric_limits::min() / static_cast(2), + std::numeric_limits::min() / static_cast(2) - static_cast(1), + static_cast(-2), + static_cast(-1), + static_cast(0), + static_cast(1), + static_cast(2), + std::numeric_limits::max() / static_cast(2) + static_cast(1), + std::numeric_limits::max() / static_cast(2), + std::numeric_limits::max() / static_cast(2) - static_cast(1), + std::numeric_limits::max() - static_cast(3), + std::numeric_limits::max() - static_cast(2), + std::numeric_limits::max() - static_cast(1), + std::numeric_limits::max(), + }; + + std::uniform_int_distribution dist( + std::numeric_limits::min(), + std::numeric_limits::max() + ); + + while (inputs.size() < n) { + inputs.push_back(dist(rng)); + } + + return inputs; +} + +template +void stress_test_checked(F op, G checked_op) { + INFO("L=" << typeid(L).name()); + INFO("R=" << typeid(R).name()); + INFO("O=" << typeid(O).name()); + + auto left_inputs = generate_inputs(std::mt19937_64 {0}, 1000); + auto right_inputs = generate_inputs(std::mt19937_64 {1}, 1000); + + for (auto a : left_inputs) { + for (auto b : right_inputs) { + O c {}; + + INFO("left=" << a); + INFO("right=" << b); + + __int128 expected = op(__int128(a), __int128(b)); + + if (expected >= __int128(std::numeric_limits::min()) + && expected <= __int128(std::numeric_limits::max())) { + INFO("expected=" << static_cast(expected)); + REQUIRE_NOTHROW(c = checked_op(a, b)); + + INFO("gotten=" << c); + REQUIRE(c == static_cast(expected)); + } else { + INFO("expected="); + REQUIRE_THROWS(checked_op(a, b)); + } + } + } +} + +template +void stress_test_checked_add() { + stress_test_checked(std::plus<__int128> {}, checked_add); +} + +TEST_CASE("checked_add (stress test)", "[.][slow]") { + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); + stress_test_checked_add(); +} + +template +void stress_test_checked_sub() { + stress_test_checked(std::minus<__int128> {}, checked_sub); +} + +TEST_CASE("checked_sub (stress test)", "[.][slow]") { + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); + stress_test_checked_sub(); +} + +template +void stress_test_checked_mul() { + stress_test_checked( + [](__int128 a, __int128 b) { + auto x = a >= 0 ? (unsigned __int128)a : (unsigned __int128)-a; + auto y = b >= 0 ? (unsigned __int128)b : (unsigned __int128)-b; + auto z = x * y; + return ((a >= 0) == (b >= 0)) ? __int128(z) : -__int128(z); + }, + checked_mul + ); +} + +TEST_CASE("checked_mul (stress test)", "[.][slow]") { + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); + stress_test_checked_mul(); +} + +template +void stress_test_checked_div() { + stress_test_checked( + [](__int128 a, __int128 b) { return b != 0 ? a / b : __int128(1) << 127; }, + checked_div + ); +} + +TEST_CASE("checked_div (stress test)", "[.][slow]") { + stress_test_checked_div(); + stress_test_checked_mul(); + stress_test_checked_div(); + stress_test_checked_div(); + + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); + stress_test_checked_div(); +} + +template +void stress_test_checked_rem() { + stress_test_checked( + [](__int128 a, __int128 b) { return b != 0 ? a % b : __int128(1) << 127; }, + checked_rem + ); +} + +TEST_CASE("checked_rem (stress test)", "[.][slow]") { + stress_test_checked_rem(); + stress_test_checked_mul(); + stress_test_checked_rem(); + stress_test_checked_rem(); + + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); + stress_test_checked_rem(); +} \ No newline at end of file diff --git a/test/core/test_const_value.cpp b/test/core/test_const_value.cpp new file mode 100644 index 00000000..c35cb357 --- /dev/null +++ b/test/core/test_const_value.cpp @@ -0,0 +1,40 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/const_value.hpp" + +#define CHECK_TYPE(V, ...) CHECK(std::is_same::value) + +using namespace kmm; + +TEST_CASE("ConstValue") { + ConstValue<5> v; + + CHECK(v.value == 5); + CHECK(static_cast(v) == 5); + CHECK(v() == 5); + + SECTION("converting constructor from a compatible ConstValue") { + ConstValue<5L> from_int(v); + CHECK(from_int.value == 5); + } + + SECTION("operators") { + CHECK_TYPE(-v, ConstValue<-5>); + CHECK_TYPE(+v, ConstValue<5>); + + ConstValue<2> a; + ConstValue<3> b; + + CHECK_TYPE(a + b, ConstValue<5>); + CHECK_TYPE(a - b, ConstValue<-1>); + CHECK_TYPE(a * b, ConstValue<6>); + CHECK_TYPE(b / a, ConstValue<1>); + + CHECK(a + int(b) == 5); + CHECK(a - int(b) == -1); + CHECK(a * int(b) == 6); + CHECK(b / int(a) == 1); + } +} diff --git a/test/core/test_distribution.cpp b/test/core/test_distribution.cpp new file mode 100644 index 00000000..232c6cba --- /dev/null +++ b/test/core/test_distribution.cpp @@ -0,0 +1,49 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/core/distribution.hpp" + +using namespace kmm; + +TEST_CASE("Distribution grid/extent/offset") { + Distribution<2> dist(Shape<2>(100, 100), Shape<2>(32, 32)); + + SECTION("evenly-divisible chunk shape yields ceil(total / chunk) grid") { + CHECK(dist.total_shape() == Shape<2>(100, 100)); + CHECK(dist.chunk_shape() == Shape<2>(32, 32)); + CHECK(dist.grid_shape() == Shape<2>(4, 4)); + CHECK(dist.num_chunks() == 16); + } + + SECTION("interior chunks have the nominal extent") { + CHECK(dist.chunk_extent(Point<2>(0, 0)) == Shape<2>(32, 32)); + CHECK(dist.chunk_extent(Point<2>(1, 2)) == Shape<2>(32, 32)); + } + + SECTION("edge chunks are clipped to what remains") { + CHECK(dist.chunk_extent(Point<2>(3, 0)) == Shape<2>(4, 32)); + CHECK(dist.chunk_extent(Point<2>(3, 3)) == Shape<2>(4, 4)); + } + + SECTION("chunk_offset is the grid index scaled by the nominal chunk shape") { + CHECK(dist.chunk_offset(Point<2>(0, 0)) == Point<2>(0, 0)); + CHECK(dist.chunk_offset(Point<2>(3, 1)) == Point<2>(96, 32)); + } +} + +TEST_CASE("Distribution::linear_index and Distribution::unravel are inverses") { + Distribution<3> dist(Shape<3>(10, 20, 30), Shape<3>(3, 7, 11)); + + for (size_t linear = 0; linear < dist.num_chunks(); linear++) { + auto grid_index = dist.unravel(linear); + CHECK(dist.linear_index(grid_index) == linear); + } +} + +TEST_CASE("Distribution handles a single chunk covering the whole domain") { + Distribution<2> dist(Shape<2>(50, 50), Shape<2>(100, 100)); + + CHECK(dist.grid_shape() == Shape<2>(1, 1)); + CHECK(dist.num_chunks() == 1); + CHECK(dist.chunk_extent(Point<2>(0, 0)) == Shape<2>(50, 50)); + CHECK(dist.chunk_offset(Point<2>(0, 0)) == Point<2>(0, 0)); +} diff --git a/test/core/test_fast_divisor.cpp b/test/core/test_fast_divisor.cpp new file mode 100644 index 00000000..af30c3c5 --- /dev/null +++ b/test/core/test_fast_divisor.cpp @@ -0,0 +1,183 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/fast_divisor.hpp" + +using namespace kmm; + +using u32 = uint32_t; +using u64 = uint64_t; + +TEST_CASE("FastDivisor") { + for (u32 divisor : {u32(1), u32(2), u32(3), u32(7), u32(64), u32(1000), u32(1u << 30)}) { + FastDivisor fd(divisor); + CAPTURE(divisor); + + CHECK(fd.get() == divisor); + + for (u32 n : {u32(0), u32(1), divisor - 1, divisor + 1, u32(12345), u32(0x7fffffff)}) { + CAPTURE(n); + CHECK(fd.divide(n) == n / divisor); + CHECK(fd.modulo(n) == n % divisor); + CHECK((n / fd) == n / divisor); + CHECK((n % fd) == n % divisor); + } + } + + CHECK(FastDivisor().divide(777) == 777); + CHECK_THROWS_AS(FastDivisor(0), std::runtime_error); +} + +TEST_CASE("FastDivisor") { + for (u64 divisor : {u64(1), u64(2), u64(3), u64(7), u64(64), u64(1000), u64(1u << 30)}) { + FastDivisor fd(divisor); + CAPTURE(divisor); + + CHECK(fd.get() == divisor); + + for (u64 n : {u64(0), u64(1), divisor - 1, divisor + 1, u64(12345), u64(0x7fffffff)}) { + CAPTURE(n); + CHECK(fd.divide(n) == n / divisor); + CHECK(fd.modulo(n) == n % divisor); + CHECK((n / fd) == n / divisor); + CHECK((n % fd) == n % divisor); + } + } + + CHECK(FastDivisor().divide(777) == 777); + CHECK_THROWS_AS(FastDivisor(0), std::runtime_error); +} + +static std::vector test_extents = { + 1, + 2, + 3, + 4, + 100, + 1337, + INT_MAX / 3, + INT_MAX / 2 - 1, + INT_MAX / 2, + INT_MAX / 2 + 1, + INT_MAX - 2, + INT_MAX - 1, + 0xc8000000, + INT_MAX +}; + +static std::vector test_indices = { + 0, + 1, + 2, + 3, + 4, + 100, + 200, + 1337, + 100000, + INT_MAX / 3, + INT_MAX / 2 - 1, + INT_MAX / 2, + INT_MAX / 2 + 1, + INT_MAX - 2, + INT_MAX - 1, + 0xc8000000, + INT_MAX +}; + +TEST_CASE("IndexMapper") { + IndexMapper mapper {}; + CHECK(mapper.ravel({}) == 0); + Vec p; + + for (uint32_t index : test_indices) { + CAPTURE(index); + CHECK(mapper.unravel(index, p) == (index == 0)); + } +} + +TEST_CASE("IndexMapper") { + for (uint32_t extent : test_extents) { + // Total volume cannot exceed max_numerator. + if (extent >= FastDivisor::max_numerator) { + continue; + } + + CAPTURE(extent); + + IndexMapper mapper {{extent}}; + Vec p {}; + + // test unravel -> ravel + for (uint32_t linear_index : test_indices) { + if (linear_index < extent) { + CHECK(mapper.unravel(linear_index, p)); + CHECK(p[0] == linear_index); + CHECK(mapper.ravel(p) == linear_index); + } else { + CHECK_FALSE(mapper.unravel(linear_index, p)); + } + } + + // test ravel -> unravel + for (uint32_t index : test_indices) { + CAPTURE(index); + + if (index < extent) { + CHECK(mapper.ravel({index}) == index); + CHECK(mapper.unravel(index, p)); + CHECK(p[0] == index); + } + } + } +} + +TEST_CASE("IndexMapper") { + for (uint32_t extent0 : test_extents) { + for (uint32_t extent1 : test_extents) { + // Total volume cannot exceed uint32_t max. + unsigned __int128 volume = (unsigned __int128)(extent0)*extent1; + + if (volume > std::numeric_limits::max()) { + continue; + } + + CAPTURE(extent0); + CAPTURE(extent1); + + IndexMapper mapper {{extent0, extent1}}; + Vec p {0, 0}; + + // test unravel -> ravel + for (uint32_t linear_index : test_indices) { + CAPTURE(linear_index); + + if (linear_index < volume) { + CHECK(mapper.unravel(linear_index, p)); + CHECK(p[0] == linear_index % extent0); + CHECK(p[1] == linear_index / extent0); + CHECK(mapper.ravel(p) == linear_index); + } else { + CHECK_FALSE(mapper.unravel(linear_index, p)); + } + } + + // test ravel -> unravel + for (uint32_t index0 : test_indices) { + for (uint32_t index1 : test_indices) { + if (index0 < extent0 && index1 < extent1) { + uint32_t linear_index = index0 + index1 * extent0; + CAPTURE(index0); + CAPTURE(index1); + + CHECK(mapper.ravel(Vec {index0, index1}) == linear_index); + CHECK(mapper.unravel(uint32_t(linear_index), p)); + CHECK(p[0] == index0); + CHECK(p[1] == index1); + } + } + } + } + } +} diff --git a/test/core/test_identifiers.cpp b/test/core/test_identifiers.cpp deleted file mode 100644 index 70f8f803..00000000 --- a/test/core/test_identifiers.cpp +++ /dev/null @@ -1,106 +0,0 @@ -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/core/identifiers.hpp" - -using namespace kmm; - -TEST_CASE("DeviceStreamSet") { - auto empty = DeviceStreamSet {}; - auto all = DeviceStreamSet::all(); - auto one = DeviceStreamSet {1}; - auto two = DeviceStreamSet {1, 5}; - auto range = DeviceStreamSet::range(4, 8); - - SECTION("contains") { - for (size_t i = 0; i < DeviceStreamSet::MAX_SIZE; i++) { - REQUIRE_FALSE(empty.contains(i)); - } - - for (size_t i = 0; i < DeviceStreamSet::MAX_SIZE; i++) { - REQUIRE(all.contains(i)); - } - - REQUIRE_FALSE(one.contains(0)); - REQUIRE(one.contains(1)); - REQUIRE_FALSE(one.contains(5)); - - REQUIRE_FALSE(two.contains(0)); - REQUIRE(two.contains(1)); - REQUIRE(two.contains(5)); - - REQUIRE_FALSE(range.contains(3)); - REQUIRE(range.contains(4)); - REQUIRE(range.contains(5)); - REQUIRE(range.contains(6)); - REQUIRE(range.contains(7)); - REQUIRE_FALSE(range.contains(8)); - } - - SECTION("contains subset") { - REQUIRE(empty.contains(empty)); - REQUIRE(all.contains(empty)); - REQUIRE(one.contains(empty)); - REQUIRE(two.contains(empty)); - - REQUIRE_FALSE(empty.contains(all)); - REQUIRE(all.contains(all)); - REQUIRE_FALSE(one.contains(all)); - REQUIRE_FALSE(two.contains(all)); - - REQUIRE_FALSE(empty.contains(one)); - REQUIRE(all.contains(one)); - REQUIRE(one.contains(one)); - REQUIRE(two.contains(one)); - - REQUIRE_FALSE(empty.contains(two)); - REQUIRE(all.contains(two)); - REQUIRE_FALSE(one.contains(two)); - REQUIRE(two.contains(two)); - } - - SECTION("operator&") { - REQUIRE((empty & empty) == empty); - REQUIRE((all & empty) == empty); - REQUIRE((one & empty) == empty); - REQUIRE((two & empty) == empty); - REQUIRE((range & empty) == empty); - - REQUIRE((empty & all) == empty); - REQUIRE((all & all) == all); - REQUIRE((one & all) == one); - REQUIRE((two & all) == two); - REQUIRE((range & all) == range); - - REQUIRE((one & one) == one); - REQUIRE((one & two) == one); - REQUIRE((one & range) == empty); - REQUIRE((two & one) == one); - REQUIRE((two & two) == two); - REQUIRE((two & range) == DeviceStreamSet {5}); - REQUIRE((range & one) == empty); - REQUIRE((range & two) == DeviceStreamSet {5}); - REQUIRE((range & range) == range); - } - - SECTION("operator==") { - REQUIRE(empty == empty); - REQUIRE(all == all); - REQUIRE(one == one); - REQUIRE(two == two); - REQUIRE(range == range); - - REQUIRE_FALSE(empty == range); - REQUIRE_FALSE(all == empty); - REQUIRE_FALSE(one == all); - REQUIRE_FALSE(two == one); - REQUIRE_FALSE(range == two); - - REQUIRE(empty != range); - REQUIRE(all != empty); - REQUIRE(one != all); - REQUIRE(two != one); - REQUIRE(range != two); - } -} \ No newline at end of file diff --git a/test/core/test_integer_fun.cpp b/test/core/test_integer_fun.cpp new file mode 100644 index 00000000..bf36a34f --- /dev/null +++ b/test/core/test_integer_fun.cpp @@ -0,0 +1,237 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/integer_fun.hpp" + +using namespace kmm; + +using i32 = signed int; +using u32 = unsigned int; + +TEST_CASE("div_floor") { + CHECK(div_floor(i32(0), i32(3)) == 0); + CHECK(div_floor(i32(10), i32(3)) == 3); + CHECK(div_floor(i32(10), i32(1)) == 10); + CHECK(div_floor(i32(10), i32(103)) == 0); + + CHECK(div_floor(i32(0), i32(-3)) == 0); + CHECK(div_floor(i32(10), i32(-3)) == -4); + CHECK(div_floor(i32(10), i32(-1)) == -10); + CHECK(div_floor(i32(10), i32(-103)) == -1); + + CHECK(div_floor(i32(-10), i32(3)) == -4); + CHECK(div_floor(i32(-10), i32(1)) == -10); + CHECK(div_floor(i32(-10), i32(103)) == -1); + CHECK(div_floor(i32(-10), i32(-3)) == 3); + CHECK(div_floor(i32(-10), i32(-1)) == 10); + CHECK(div_floor(i32(-10), i32(-103)) == 0); + + CHECK(div_floor(u32(0), u32(3)) == 0); + CHECK(div_floor(u32(10), u32(3)) == 3); + CHECK(div_floor(u32(10), u32(103)) == 0); + + CHECK(div_floor(INT_MIN, INT_MIN) == 1); + CHECK(div_floor(INT_MAX, INT_MIN) == -1); + CHECK(div_floor(INT_MIN, INT_MAX) == -2); + CHECK(div_floor(INT_MAX, INT_MAX) == 1); + + CHECK(div_floor(INT_MIN, 5) == INT_MIN / 5 - 1); + CHECK(div_floor(INT_MAX, 5) == INT_MAX / 5); + CHECK(div_floor(5, INT_MIN) == -1); + CHECK(div_floor(5, INT_MAX) == 0); + + CHECK(div_floor(u32(0), UINT_MAX) == 0); + CHECK(div_floor(UINT_MAX, UINT_MAX) == 1); + CHECK(div_floor(UINT_MAX, u32(5)) == UINT_MAX / 5); + CHECK(div_floor(u32(5), UINT_MAX) == 0); +} + +TEST_CASE("div_ceil") { + CHECK(div_ceil(i32(0), i32(3)) == 0); + CHECK(div_ceil(i32(10), i32(3)) == 4); + CHECK(div_ceil(i32(10), i32(1)) == 10); + CHECK(div_ceil(i32(10), i32(103)) == 1); + + CHECK(div_ceil(i32(0), i32(-3)) == 0); + CHECK(div_ceil(i32(10), i32(-3)) == -3); + CHECK(div_ceil(i32(10), i32(-1)) == -10); + CHECK(div_ceil(i32(10), i32(-103)) == 0); + + CHECK(div_ceil(i32(-10), i32(3)) == -3); + CHECK(div_ceil(i32(-10), i32(1)) == -10); + CHECK(div_ceil(i32(-10), i32(103)) == 0); + CHECK(div_ceil(i32(-10), i32(-3)) == 4); + CHECK(div_ceil(i32(-10), i32(-1)) == 10); + CHECK(div_ceil(i32(-10), i32(-103)) == 1); + + CHECK(div_ceil(u32(0), u32(3)) == 0); + CHECK(div_ceil(u32(10), u32(3)) == 4); + CHECK(div_ceil(u32(10), u32(103)) == 1); + + CHECK(div_ceil(INT_MIN, INT_MIN) == 1); + CHECK(div_ceil(INT_MAX, INT_MIN) == 0); + CHECK(div_ceil(INT_MIN, INT_MAX) == -1); + CHECK(div_ceil(INT_MAX, INT_MAX) == 1); + + CHECK(div_ceil(INT_MIN, 5) == INT_MIN / 5); + CHECK(div_ceil(INT_MAX, 5) == INT_MAX / 5 + 1); + CHECK(div_ceil(5, INT_MIN) == 0); + CHECK(div_ceil(5, INT_MAX) == 1); + + CHECK(div_ceil(u32(0), UINT_MAX) == 0); + CHECK(div_ceil(UINT_MAX, UINT_MAX) == 1); + CHECK(div_ceil(UINT_MAX, u32(5)) == UINT_MAX / 5); + CHECK(div_ceil(u32(5), UINT_MAX) == 1); +} + +TEST_CASE("round_up_to_multiple") { + CHECK(round_up_to_multiple(i32(5), i32(3)) == 6); + CHECK(round_up_to_multiple(i32(5), i32(-3)) == 6); + CHECK(round_up_to_multiple(i32(-5), i32(3)) == -3); + CHECK(round_up_to_multiple(i32(-5), i32(-3)) == -3); + + CHECK(round_up_to_multiple(i32(0), i32(3)) == 0); + CHECK(round_up_to_multiple(i32(9), i32(3)) == 9); + CHECK(round_up_to_multiple(i32(10), i32(3)) == 12); + CHECK(round_up_to_multiple(i32(10), i32(1)) == 10); + CHECK(round_up_to_multiple(i32(10), i32(103)) == 103); + + CHECK(round_up_to_multiple(i32(-9), i32(3)) == -9); + CHECK(round_up_to_multiple(i32(-10), i32(3)) == -9); + CHECK(round_up_to_multiple(i32(-1), i32(3)) == 0); + CHECK(round_up_to_multiple(i32(-12), i32(4)) == -12); + CHECK(round_up_to_multiple(i32(-13), i32(4)) == -12); + + CHECK(round_up_to_multiple(i32(0), i32(-3)) == 0); + CHECK(round_up_to_multiple(i32(9), i32(-3)) == 9); + CHECK(round_up_to_multiple(i32(10), i32(-3)) == 12); + CHECK(round_up_to_multiple(i32(10), i32(-1)) == 10); + CHECK(round_up_to_multiple(i32(10), i32(-103)) == 103); + + CHECK(round_up_to_multiple(i32(-9), i32(-3)) == -9); + CHECK(round_up_to_multiple(i32(-10), i32(-3)) == -9); + CHECK(round_up_to_multiple(i32(-1), i32(-3)) == 0); + CHECK(round_up_to_multiple(i32(-12), i32(-4)) == -12); + CHECK(round_up_to_multiple(i32(-13), i32(-4)) == -12); + + CHECK(round_up_to_multiple(u32(0), u32(3)) == 0); + CHECK(round_up_to_multiple(u32(9), u32(3)) == 9); + CHECK(round_up_to_multiple(u32(10), u32(3)) == 12); + CHECK(round_up_to_multiple(u32(10), u32(1)) == 10); + CHECK(round_up_to_multiple(u32(10), u32(103)) == 103); + + CHECK(round_up_to_multiple(INT_MIN, INT_MAX) == INT_MIN + 1); + CHECK(round_up_to_multiple(INT_MIN + 1, INT_MAX) == INT_MIN + 1); + CHECK(round_up_to_multiple(-1, INT_MAX) == 0); + CHECK(round_up_to_multiple(0, INT_MAX) == 0); + CHECK(round_up_to_multiple(1, INT_MAX) == INT_MAX); + CHECK(round_up_to_multiple(INT_MAX - 1, INT_MAX) == INT_MAX); + CHECK(round_up_to_multiple(INT_MAX, INT_MAX) == INT_MAX); + + CHECK(round_up_to_multiple(INT_MIN, INT_MIN) == INT_MIN); + CHECK(round_up_to_multiple(INT_MIN + 1, INT_MIN) == 0); + CHECK(round_up_to_multiple(-1, INT_MAX) == 0); + CHECK(round_up_to_multiple(0, INT_MIN) == 0); + + CHECK(round_up_to_multiple(u32(0), UINT_MAX) == 0); + CHECK(round_up_to_multiple(u32(1), UINT_MAX) == UINT_MAX); + CHECK(round_up_to_multiple(UINT_MAX - u32(1), UINT_MAX) == UINT_MAX); + CHECK(round_up_to_multiple(UINT_MAX, UINT_MAX) == UINT_MAX); + + // -1 * INT_MIN >= INT_MAX, so there will overflow + CHECK_THROWS(round_up_to_multiple(1, INT_MIN)); + CHECK_THROWS(round_up_to_multiple(INT_MAX - 1, INT_MIN)); + CHECK_THROWS(round_up_to_multiple(INT_MAX, INT_MIN)); + + // INT_MAX/UINT_MAX are odd, so rounding up to a multiple of 2 overflows. + CHECK_THROWS(round_up_to_multiple(INT_MAX, 2)); + CHECK_THROWS(round_up_to_multiple(UINT_MAX, u32(2))); +} + +TEST_CASE("unsigned_abs") { + CHECK(unsigned_abs(i32(0)) == u32(0)); + CHECK(unsigned_abs(i32(5)) == u32(5)); + CHECK(unsigned_abs(i32(-5)) == u32(5)); + CHECK(unsigned_abs(INT_MAX) == u32(INT_MAX)); + CHECK(unsigned_abs(INT_MIN) == u32(INT_MAX) + u32(1)); + + CHECK(unsigned_abs(u32(0)) == u32(0)); + CHECK(unsigned_abs(u32(5)) == u32(5)); + CHECK(unsigned_abs(UINT_MAX) == UINT_MAX); + + STATIC_REQUIRE(std::is_same_v); +} + +TEST_CASE("is_divisible") { + CHECK(is_divisible(10, 5)); + CHECK(is_divisible(10, 2)); + CHECK_FALSE(is_divisible(10, 3)); + + CHECK(is_divisible(0, 5)); + CHECK(is_divisible(-10, 5)); + CHECK(is_divisible(10, -5)); + CHECK_FALSE(is_divisible(-10, 3)); + + // division by zero yields false rather than throwing + CHECK_FALSE(is_divisible(10, 0)); + CHECK_FALSE(is_divisible(0, 0)); + + // mixed signed/unsigned operands + CHECK(is_divisible(10, 5u)); + CHECK(is_divisible(-10, 5u)); + CHECK_FALSE(is_divisible(-10, 3u)); + + // large unsigned operands whose remainder does not fit in int64_t + CHECK(is_divisible(ULONG_MAX, ULONG_MAX)); + CHECK_FALSE(is_divisible(ULONG_MAX, ULONG_MAX - 1)); +} + +TEST_CASE("is_power_of_two") { + CHECK_FALSE(is_power_of_two(INT_MIN)); + CHECK_FALSE(is_power_of_two(-1)); + CHECK_FALSE(is_power_of_two(0)); + CHECK(is_power_of_two(1)); + CHECK(is_power_of_two(2)); + CHECK_FALSE(is_power_of_two(3)); + CHECK(is_power_of_two(4)); + CHECK_FALSE(is_power_of_two(5)); + CHECK(is_power_of_two(128)); + CHECK_FALSE(is_power_of_two(100)); + CHECK(is_power_of_two(1 << 30)); + CHECK_FALSE(is_power_of_two(INT_MAX)); + + CHECK_FALSE(is_power_of_two(u32(0))); + CHECK(is_power_of_two(u32(1))); + CHECK(is_power_of_two(u32(2))); + CHECK_FALSE(is_power_of_two(u32(100))); + CHECK(is_power_of_two(u32(INT_MAX) + u32(1))); + CHECK_FALSE(is_power_of_two(UINT_MAX)); +} + +TEST_CASE("round_up_to_power_of_two") { + // <1 always becomes 1 + CHECK(round_up_to_power_of_two(INT_MIN) == 1); + CHECK(round_up_to_power_of_two(-1) == 1); + CHECK(round_up_to_power_of_two(0) == 1); + + CHECK(round_up_to_power_of_two(5) == 8); + CHECK(round_up_to_power_of_two(100) == 128); + CHECK(round_up_to_power_of_two(128) == 128); + CHECK(round_up_to_power_of_two(1000) == 1024); + + CHECK(round_up_to_power_of_two(u32(0)) == 1); + CHECK(round_up_to_power_of_two(u32(5)) == 8); + CHECK(round_up_to_power_of_two(u32(100)) == 128); + CHECK(round_up_to_power_of_two(u32(128)) == 128); + CHECK(round_up_to_power_of_two(u32(1000)) == 1024); + CHECK(round_up_to_power_of_two(u32(INT_MAX)) == u32(INT_MAX) + u32(1)); + CHECK(round_up_to_power_of_two(u32(INT_MAX) + u32(1)) == u32(INT_MAX) + u32(1)); + + // These should overflow + CHECK_THROWS(round_up_to_power_of_two(INT_MAX)); + CHECK_THROWS(round_up_to_power_of_two(u32(INT_MAX) + u32(2))); + CHECK_THROWS(round_up_to_power_of_two(UINT_MAX - 1)); + CHECK_THROWS(round_up_to_power_of_two(UINT_MAX)); +} \ No newline at end of file diff --git a/test/core/test_key_value.cpp b/test/core/test_key_value.cpp new file mode 100644 index 00000000..65eba228 --- /dev/null +++ b/test/core/test_key_value.cpp @@ -0,0 +1,57 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/core/key_value.hpp" + +using namespace kmm; + +TEST_CASE("KeyValue construction") { + // default constructor + KeyValue empty; + (void)empty; + + // value constructor + KeyValue kv(42, 100); + CHECK(kv.key == 42); + CHECK(kv.value == 100); +} + +TEST_CASE("KeyValue::operator==/!=") { + KeyValue a(1, 10); + KeyValue b(1, 10); + KeyValue c(2, 10); + KeyValue d(1, 20); + + CHECK(a == b); + CHECK_FALSE(a != b); + + CHECK(a != c); + CHECK_FALSE(a == c); + + CHECK(a != d); + CHECK_FALSE(a == d); +} + +TEST_CASE("KeyValue::operator=/>") { + // low < high + KeyValue low_value(5, 1); + KeyValue high_value(1, 2); + CHECK(low_value < high_value); + CHECK(low_value <= high_value); + CHECK(high_value > low_value); + CHECK(high_value >= low_value); + + KeyValue a(1, 10); + KeyValue b(2, 10); + + // ties are broken using key + CHECK(a < b); + CHECK(a <= b); + CHECK(b > a); + CHECK(b >= a); + + // two equal items + CHECK_FALSE(a < a); + CHECK_FALSE(a > a); + CHECK(a <= a); + CHECK(a >= a); +} diff --git a/test/core/test_layout.cpp b/test/core/test_layout.cpp new file mode 100644 index 00000000..6458b3a8 --- /dev/null +++ b/test/core/test_layout.cpp @@ -0,0 +1,512 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/core/layout.hpp" + +using namespace kmm; + +TEST_CASE("Layout basics") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(TestLayout::rank == 2); + CHECK(layout.domain() == Shape<2, int>(4, 5)); + CHECK(layout.mapping().get(ConstIndex<0>()) == 5); +} + +TEST_CASE("Layout::extent/begin/end/origin") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(layout.extent(0) == 4); + CHECK(layout.extent(1) == 5); + CHECK(layout.extent(5) == 1); // out-of-range axis defaults to one + + CHECK(layout.begin(0) == 0); + CHECK(layout.end(0) == 4); + CHECK(layout.begin(1) == 0); + CHECK(layout.end(1) == 5); + CHECK(layout.origin(0) == 0); + CHECK(layout.origin(1) == 0); + + CHECK(layout.begin() == Point<2, int>(0, 0)); + CHECK(layout.end() == Point<2, int>(4, 5)); +} + +TEST_CASE("Layout::shape/bounds/size/is_empty") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(layout.shape() == Shape<2, int>(4, 5)); + CHECK(layout.bounds().shape() == Shape<2, int>(4, 5)); + CHECK(layout.size() == 20); + CHECK_FALSE(layout.is_empty()); + + SECTION("an axis of size zero makes the layout empty") { + TestLayout empty_layout(Shape<2, int>(0, 5), {5, 1}); + CHECK(empty_layout.is_empty()); + CHECK(empty_layout.size() == 0); + } +} + +TEST_CASE("Layout::contains") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(layout.contains(Point<2, int>(0, 0))); + CHECK(layout.contains(Point<2, int>(3, 4))); + CHECK_FALSE(layout.contains(Point<2, int>(4, 0))); + CHECK_FALSE(layout.contains(Point<2, int>(0, 5))); +} + +TEST_CASE("Layout::stride/strides") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(layout.stride(0) == 5); + CHECK(layout.stride(1) == 1); + + auto strides = layout.strides(); + CHECK(strides[0] == 5); + CHECK(strides[1] == 1); +} + +TEST_CASE("Layout::local_offset/offset") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + CHECK(layout.local_offset(Point<2, int>(0, 0)) == 0); + CHECK(layout.local_offset(Point<2, int>(1, 2)) == 5 * 1 + 1 * 2); + CHECK(layout.local_offset(Point<2, int>(3, 4)) == 5 * 3 + 1 * 4); + + SECTION("offset combines base_offset and local_offset") { + using BoundsLayout = ::kmm::Layout, Strided>; + BoundsLayout bl(Bounds<2, int>(Range(2, 6), Range(3, 8)), {5, 1}); + + CHECK(bl.offset(Point<2, int>(2, 3)) == 0); // offset at the domain's own origin is zero + CHECK(bl.offset(Point<2, int>(3, 5)) == 7); + } +} + +TEST_CASE("Layout::is_contiguous") { + SECTION("row-major strides are contiguous in row-major order, not column-major") { + using TestLayout = ::kmm::Layout, Strides>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + CHECK(layout.is_contiguous()); + CHECK(layout.is_contiguous(MemoryOrder::RowMajor)); + CHECK_FALSE(layout.is_contiguous(MemoryOrder::ColMajor)); + } + + SECTION("column-major layout is contiguous in column-major order, not row-major") { + auto layout = make_layout(Shape<3, int>(4, 5, 6)); + + CHECK(layout.is_contiguous(MemoryOrder::ColMajor)); + CHECK_FALSE(layout.is_contiguous(MemoryOrder::RowMajor)); + CHECK_FALSE(layout.is_contiguous()); // default order is row-major + } + + SECTION("padded strides are not contiguous") { + auto layout = make_layout>(Shape<2, int>(5, 3)); + + CHECK_FALSE(layout.is_contiguous()); + } + + SECTION("rank 0 layout is always contiguous") { + auto layout = make_layout(Shape<0, int>()); + + CHECK(layout.is_contiguous()); + CHECK(layout.is_contiguous(MemoryOrder::ColMajor)); + } +} + +TEST_CASE("Layout::is_mapping_from_policy") { + auto layout = make_layout(Shape<2, int>(4, 5)); + CHECK(layout.is_mapping_from_policy()); + CHECK_FALSE(layout.is_mapping_from_policy()); + + using TestLayout = ::kmm::Layout, Strided>; + TestLayout mismatched(Shape<2, int>(4, 5), TestLayout::mapping_type(1, 1)); + CHECK_FALSE(mismatched.is_mapping_from_policy()); +} + +TEST_CASE("Layout::with_domain/with_mapping") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<2, int>(4, 5), {5, 1}); + + auto same_shape = layout.with_domain(Shape<2, int>(2, 3)); + CHECK(same_shape.size() == 6); + + auto new_strides = layout.with_mapping(Strides(100, 10)); + CHECK(new_strides.stride(0) == 100); + CHECK(new_strides.stride(1) == 10); + CHECK(new_strides.shape() == Shape<2, int>(4, 5)); +} + +TEST_CASE("Layout::zero_origin") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Bounds<2, int>(Range(2, 6), Range(3, 8)), {5, 1}); + CHECK(layout.base_offset() == -(2 * 5 + 3)); + + auto zo = layout.zero_origin(); + CHECK(zo.begin() == Point<2, int>(0, 0)); + CHECK(zo.shape() == Shape<2, int>(4, 5)); + CHECK(zo.base_offset() == 0); +} + +TEST_CASE("Layout::move_origin") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + auto moved = layout.move_origin(Point<3, int>(1, 2, 3)); + + CHECK(moved.begin() == Point<3, int>(0, 0, 0)); + CHECK(moved.shape() == Shape<3, int>(4, 5, 6)); + CHECK(moved.size() == 120); + CHECK(moved.stride(0) == 30); + CHECK(moved.stride(1) == 6); + CHECK(moved.stride(2) == 1); + CHECK(moved.local_offset({1, 2, 3}) == 1 * 30 + 2 * 6 + 3 * 1); + CHECK(moved.base_offset() == 1 * 30 + 2 * 6 + 3); +} + +TEST_CASE("Layout::restrict_bounds") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + auto restricted = layout.restrict_bounds( + Bounds<3, int>(Range(1, 3), Range(1, 4), Range(0, 6)) + ); + + CHECK(restricted.begin() == Point<3, int>(1, 1, 0)); + CHECK(restricted.end() == Point<3, int>(3, 4, 6)); + CHECK(restricted.size() == 2 * 3 * 6); + CHECK(restricted.base_offset() == 0); +} + +TEST_CASE("Layout::restrict_axis") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + auto restricted = layout.restrict_axis<0>(1, 3); + + CHECK(restricted.extent(0) == 2); + CHECK(restricted.extent(1) == 5); + CHECK(restricted.extent(2) == 6); + CHECK(restricted.begin(0) == 1); + CHECK(restricted.end(0) == 3); + CHECK(restricted.base_offset() == 0); +} + +TEST_CASE("Layout::slice_bounds") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + auto sliced = + layout.slice_bounds(Bounds<3, int>(Range(1, 3), Range(1, 4), Range(0, 6))); + + CHECK(sliced.begin() == Point<3, int>(0, 0, 0)); + CHECK(sliced.extent(0) == 2); + CHECK(sliced.extent(1) == 3); + CHECK(sliced.extent(2) == 6); + CHECK(sliced.base_offset() == 1 * 30 + 1 * 6); +} + +TEST_CASE("Layout::drop_axis") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + SECTION("dropping the middle axis keeps the outer two, in order") { + auto dropped = layout.drop_axis<1>(2); + + CHECK(dropped.extent(0) == 4); + CHECK(dropped.extent(1) == 6); + CHECK(dropped.size() == 24); + CHECK(dropped.stride(0) == 30); + CHECK(dropped.stride(1) == 1); + CHECK(dropped.base_offset() == 2 * 6); + } + + SECTION("dropping the first axis") { + auto dropped = layout.drop_axis<0>(0); + + CHECK(dropped.extent(0) == 5); + CHECK(dropped.extent(1) == 6); + CHECK(dropped.stride(0) == 6); + CHECK(dropped.stride(1) == 1); + CHECK(dropped.base_offset() == 0); + } + + SECTION("dropping the last axis") { + auto dropped = layout.drop_axis<2>(0); + + CHECK(dropped.extent(0) == 4); + CHECK(dropped.extent(1) == 5); + CHECK(dropped.stride(0) == 30); + CHECK(dropped.stride(1) == 6); + CHECK(dropped.base_offset() == 0); + } +} + +TEST_CASE("Layout::insert_axis") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + SECTION("inserting in the middle shifts later axes and gets stride zero") { + auto inserted = layout.insert_axis<1>(7); + + CHECK(inserted.rank == 4); + CHECK(inserted.extent(0) == 4); + CHECK(inserted.extent(1) == 7); + CHECK(inserted.extent(2) == 5); + CHECK(inserted.extent(3) == 6); + CHECK(inserted.stride(0) == 30); + CHECK(inserted.stride(1) == 0); + CHECK(inserted.stride(2) == 6); + CHECK(inserted.stride(3) == 1); + CHECK(inserted.size() == 4 * 7 * 5 * 6); + CHECK(inserted.base_offset() == 0); + } + + SECTION("inserting at the front") { + auto inserted = layout.insert_axis<0>(2); + + CHECK(inserted.extent(0) == 2); + CHECK(inserted.extent(1) == 4); + CHECK(inserted.extent(2) == 5); + CHECK(inserted.extent(3) == 6); + CHECK(inserted.stride(0) == 0); + CHECK(inserted.stride(1) == 30); + } + + SECTION("inserting at the end") { + auto inserted = layout.insert_axis<3>(2); + + CHECK(inserted.extent(3) == 2); + CHECK(inserted.stride(0) == 30); + CHECK(inserted.stride(3) == 0); + } + + SECTION("default extent is one") { + auto inserted = layout.insert_axis<0>(); + + CHECK(inserted.extent(0) == 1); + CHECK(inserted.stride(0) == 0); + } + + SECTION("index along the broadcast axis does not affect the linear offset") { + auto inserted = layout.insert_axis<1>(7); + + CHECK(inserted.local_offset({1, 0, 2, 3}) == 1 * 30 + 2 * 6 + 3 * 1); + CHECK(inserted.local_offset({1, 5, 2, 3}) == 1 * 30 + 2 * 6 + 3 * 1); + } +} + +TEST_CASE("Layout::reverse_axes") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + auto reversed = layout.reverse_axes(); + + CHECK(reversed.extent(0) == 6); + CHECK(reversed.extent(1) == 5); + CHECK(reversed.extent(2) == 4); + CHECK(reversed.stride(0) == 1); + CHECK(reversed.stride(1) == 6); + CHECK(reversed.stride(2) == 30); + CHECK(reversed.size() == 120); + CHECK(reversed.base_offset() == 0); +} + +TEST_CASE("Layout::slice_axis") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + SECTION("with a plain index drops the axis") { + auto sliced = layout.slice_axis<1>(2); + + CHECK(sliced.extent(0) == 4); + CHECK(sliced.extent(1) == 6); + CHECK(sliced.stride(0) == 30); + CHECK(sliced.stride(1) == 1); + CHECK(sliced.base_offset() == 2 * 6); + } + + SECTION("with a Range restricts the axis but keeps the rank") { + auto sliced = layout.slice_axis<0>(Range(1, 3)); + + CHECK(sliced.extent(0) == 2); + CHECK(sliced.extent(1) == 5); + CHECK(sliced.extent(2) == 6); + CHECK(sliced.begin(0) == 0); + CHECK(sliced.local_offset({1, 0, 0}) == 1 * 30); + CHECK(sliced.base_offset() == 1 * 30); + } + + SECTION("with all leaves the axis unchanged") { + auto sliced = layout.slice_axis<0>(all); + + CHECK(sliced.extent(0) == 4); + CHECK(sliced.extent(1) == 5); + CHECK(sliced.extent(2) == 6); + CHECK(sliced.base_offset() == 0); + } + + SECTION("with new_axis inserts a broadcast axis") { + auto sliced = layout.slice_axis<1>(new_axis); + + CHECK(sliced.rank == 4); + CHECK(sliced.extent(0) == 4); + CHECK(sliced.extent(1) == 1); + CHECK(sliced.extent(2) == 5); + CHECK(sliced.extent(3) == 6); + CHECK(sliced.stride(0) == 30); + CHECK(sliced.stride(1) == 0); + CHECK(sliced.stride(2) == 6); + CHECK(sliced.stride(3) == 1); + CHECK(sliced.base_offset() == 0); + } +} + +template +__attribute__((noinline)) auto foobar(const I& input) { + return input.slice(1, all, Range(2, 4)); +} + +TEST_CASE("Layout::slice") { + using TestLayout = ::kmm::Layout, Strided>; + TestLayout layout(Shape<3, int>(4, 5, 6), {30, 6, 1}); + + SECTION("mixing an index, all, and a Range across axes") { + foobar(layout); + auto sliced = layout.slice(1, all, Range(2, 4)); + + CHECK(sliced.extent(0) == 5); + CHECK(sliced.extent(1) == 2); + CHECK(sliced.base_offset() == 1 * 30 + 1 * 2); + } + + SECTION("all indices drops every axis") { + auto sliced = layout.slice(1, 2, 3); + + CHECK(sliced.rank == 0); + CHECK(sliced.size() == 1); + CHECK(sliced.base_offset() == 1 * 30 + 2 * 6 + 3 * 1); + } + + SECTION("all all_t leaves the layout unchanged") { + auto sliced = layout.slice(all, all, all); + + CHECK(sliced.extent(0) == 4); + CHECK(sliced.extent(1) == 5); + CHECK(sliced.extent(2) == 6); + CHECK(sliced.base_offset() == 0); + } + + SECTION("mixing new_axis with other tokens") { + auto sliced = layout.slice(all, new_axis, 2, all); + + CHECK(sliced.rank == 3); + CHECK(sliced.extent(0) == 4); + CHECK(sliced.extent(1) == 1); + CHECK(sliced.extent(2) == 6); + CHECK(sliced.stride(0) == 30); + CHECK(sliced.stride(1) == 0); + CHECK(sliced.stride(2) == 1); + CHECK(sliced.base_offset() == 2 * 6); + } +} + +TEST_CASE("make_layout") { + SECTION("rank 3") { + auto layout = make_layout(Shape<3, int>(4, 5, 6)); + + CHECK(layout.extent(0) == 4); + CHECK(layout.extent(1) == 5); + CHECK(layout.extent(2) == 6); + CHECK(layout.stride(0) == 30); + CHECK(layout.stride(1) == 6); + CHECK(layout.stride(2) == 1); + CHECK(layout.size() == 120); + CHECK(layout.local_offset(Point<3, int>(1, 2, 3)) == 1 * 30 + 2 * 6 + 3 * 1); + } + + SECTION("rank 4") { + auto layout = make_layout(Shape<4, int>(2, 3, 4, 5)); + + CHECK(layout.stride(0) == 60); + CHECK(layout.stride(1) == 20); + CHECK(layout.stride(2) == 5); + CHECK(layout.stride(3) == 1); + } + + SECTION("rank 1") { + auto layout = make_layout(Shape<1, int>(7)); + + CHECK(layout.stride(0) == 1); + CHECK(layout.size() == 7); + } + + SECTION("rank 0") { + auto layout = make_layout(Shape<0, int>()); + + CHECK(layout.size() == 1); + } + + SECTION("policy can be deduced from an argument instead of specified explicitly") { + auto layout = make_layout(Shape<2, int>(4, 5), RowMajor {}); + + CHECK(layout.stride(0) == 5); + CHECK(layout.stride(1) == 1); + } +} + +TEST_CASE("make_layout") { + SECTION("rank 0") { + auto layout = make_layout(Shape<0, int>()); + + CHECK(layout.size() == 1); + } + + SECTION("rank 1") { + auto layout = make_layout(Shape<1, int>(7)); + + CHECK(layout.stride(0) == 1); + CHECK(layout.size() == 7); + } + + SECTION("rank 3") { + auto layout = make_layout(Shape<3, int>(4, 5, 6)); + + CHECK(layout.extent(0) == 4); + CHECK(layout.extent(1) == 5); + CHECK(layout.extent(2) == 6); + CHECK(layout.stride(0) == 1); + CHECK(layout.stride(1) == 4); + CHECK(layout.stride(2) == 20); + CHECK(layout.size() == 120); + CHECK(layout.local_offset(Point<3, int>(1, 2, 3)) == 1 * 1 + 2 * 4 + 3 * 20); + } + + SECTION("rank 4") { + auto layout = make_layout(Shape<4, int>(2, 3, 4, 5)); + + CHECK(layout.stride(0) == 1); + CHECK(layout.stride(1) == 2); + CHECK(layout.stride(2) == 6); + CHECK(layout.stride(3) == 24); + } +} + +TEST_CASE("make_layout") { + auto x = make_layout(shape(5, 3), MemoryOrder::RowMajor); + CHECK(x.strides() == Vec(3, 1)); + + auto y = make_layout(shape(5, 3), MemoryOrder::ColMajor); + CHECK(y.strides() == Vec(1, 5)); + + auto xp = make_layout>(shape(5, 3), MemoryOrder::RowMajor); + CHECK(xp.strides() == Vec(4, 1)); + + auto yp = make_layout>(shape(5, 3), MemoryOrder::ColMajor); + CHECK(yp.strides() == Vec(1, 6)); +} diff --git a/test/core/test_panic.cpp b/test/core/test_panic.cpp new file mode 100644 index 00000000..ad263df5 --- /dev/null +++ b/test/core/test_panic.cpp @@ -0,0 +1,15 @@ +#include +#include +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/panic.hpp" + +using namespace kmm; + +TEST_CASE("KMM_PANIC") { + // No way to test for abort using Catch2 + // CHECK(KMM_PANIC("test panic")); +} diff --git a/test/core/test_point.cpp b/test/core/test_point.cpp new file mode 100644 index 00000000..369abf24 --- /dev/null +++ b/test/core/test_point.cpp @@ -0,0 +1,131 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/point.hpp" + +using namespace kmm; + +TEST_CASE("Point construction") { + SECTION("default constructor") { + Point<3> p; + CHECK(p[0] == 0); + CHECK(p[1] == 0); + CHECK(p[2] == 0); + } + + SECTION("variadic constructor") { + Point<3> p(1, 2, 3); + CHECK(p[0] == 1); + CHECK(p[1] == 2); + CHECK(p[2] == 3); + } + + SECTION("Vec constructor") { + Vec v {4, 5}; + Point<2> p(v); + CHECK(p[0] == 4); + CHECK(p[1] == 5); + } + + SECTION("deduction guide") { + auto p = Point(1, 2, 3, 4); + static_assert(std::is_same_v>); + CHECK(p[3] == 4); + } + + SECTION("point() function") { + auto p = point(1, 2, 3); + CHECK(p == Point<3, int>(1, 2, 3)); + } +} + +TEST_CASE("Point::one/zero") { + CHECK(Point<3>::one() == Point<3>(1, 1, 1)); + CHECK(Point<3>::zero() == Point<3>(0, 0, 0)); +} + +TEST_CASE("Point::from") { + SECTION("same dimensionality copies values") { + Vec v {1, 2, 3}; + CHECK(Point<3, int>::from(v) == Point<3>(1, 2, 3)); + } + + SECTION("growing dimensionality pads with zero") { + Vec v {1, 2}; + CHECK(Point<4, int>::from(v) == Point<4>(1, 2, 0, 0)); + } + + SECTION("shrinking dimensionality truncates") { + Vec v {1, 2, 3, 4}; + CHECK(Point<2, int>::from(v) == Point<2>(1, 2)); + } +} + +TEST_CASE("Point conversion") { + SECTION("same dimensionality, widening type") { + Point<2, int> src(1, 2); + Point<2, long> dst(src); + CHECK(dst == Point<2, long>(1, 2)); + } + + SECTION("smaller dimensionality") { + Point<3, int> src0(1, 2, 0); + Point<2, int> dst(src0); + CHECK(dst == Point<2, int>(1, 2)); + + Point<3, int> src1(1, 2, 3); + CHECK_THROWS_AS((Point<2, int>(src1)), std::overflow_error); + } + + SECTION("narrow type") { + Point<2, int> src0(300, 1); + CHECK_THROWS_AS((Point<2, signed char>(src0)), std::overflow_error); + + Point<2, int> src1(1, 2); + Point<2, signed char> dst(src1); + CHECK(dst[0] == 1); + CHECK(dst[1] == 2); + } +} + +TEST_CASE("Point::is_convertible_to") { + CHECK(Point<2, int>(1, 2).is_convertible_to<2, long>()); + CHECK_FALSE(Point<2, int>(300, 1).is_convertible_to<2, signed char>()); + CHECK(Point<3, int>(1, 2, 0).is_convertible_to<2, int>()); + CHECK_FALSE(Point<3, int>(1, 2, 3).is_convertible_to<2, int>()); +} + +TEST_CASE("Point::get_or_default") { + Point<2> p(3, 4); + CHECK(p.get_or_default(0) == 3); + CHECK(p.get_or_default(1) == 4); + CHECK(p.get_or_default(2) == 0); + CHECK(p.get_or_default(2, 99) == 99); + + Point<0> empty; + CHECK(empty.get_or_default(0) == 0); + CHECK(empty.get_or_default(0, 42) == 42); +} + +TEST_CASE("operator==(Point, Point") { + CHECK(Point<2>(1, 2) == Point<2>(1, 2)); + CHECK(Point<2>(1, 2) != Point<2>(1, 3)); + + // different N + CHECK(Point<2>(1, 2) == Point<3>(1, 2, 0)); + CHECK(Point<2>(1, 2) != Point<3>(1, 2, 3)); + + // different N and T + CHECK(Point<2>(1, 2) == Point<3, short>(short(1), short(2), short(0))); + CHECK(Point<2>(1, 2) != Point<3, short>(short(1), short(2), short(3))); +} + +TEST_CASE("concat(Point, Point)") { + Point<2, int> a(1, 2); + Point<3, int> b(3, 4, 5); + + Point<5, int> c = concat(a, b); + CHECK(c == Point<5, int>(1, 2, 3, 4, 5)); +} diff --git a/test/core/test_range.cpp b/test/core/test_range.cpp new file mode 100644 index 00000000..2702f6e2 --- /dev/null +++ b/test/core/test_range.cpp @@ -0,0 +1,189 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/range.hpp" + +using namespace kmm; + +TEST_CASE("Range construction") { + SECTION("default") { + Range r; + CHECK(r.start == 0); + CHECK(r.stop == 0); + CHECK(r.is_empty()); + } + + SECTION("single argument") { + Range r(5); + CHECK(r.start == 0); + CHECK(r.stop == 5); + CHECK_FALSE(r.is_empty()); + } + + SECTION("two arguments") { + Range r(2, 5); + CHECK(r.start == 2); + CHECK(r.stop == 5); + } + + SECTION("range functions") { + CHECK(range(5) == Range(0, 5)); + CHECK(range(2, 5) == Range(2, 5)); + } + + SECTION("one() is the range 0...1") { + CHECK(Range::one() == Range(0, 1)); + } +} + +TEST_CASE("Range::is_empty") { + CHECK_FALSE(Range(2, 5).is_empty()); + CHECK(Range(5, 5).is_empty()); + CHECK(Range(5, 2).is_empty()); +} + +TEST_CASE("Range::size") { + CHECK(Range(2, 5).size() == 3); + CHECK(Range(5, 5).size() == 0); + CHECK(Range(5, 2).size() == 0); +} + +TEST_CASE("Range iteration") { + SECTION("valid range") { + std::vector seen; + for (auto i : Range(5, 10)) { + seen.push_back(i); + } + + CHECK(seen == std::vector {5, 6, 7, 8, 9}); + } + + SECTION("empty range") { + std::vector seen; + for (auto i : Range(5, 5)) { + seen.push_back(i); + } + + CHECK(seen.empty()); + } + + SECTION("invalid range") { + std::vector seen; + for (auto i : Range(5, 2)) { + seen.push_back(i); + } + + CHECK(seen.empty()); + } +} + +TEST_CASE("Range::contains(index)") { + // valid range + Range r(5, 10); + CHECK_FALSE(r.contains(4)); + CHECK(r.contains(5)); + CHECK(r.contains(7)); + CHECK(r.contains(9)); + CHECK_FALSE(r.contains(10)); + + // invalid range + Range s(10, 5); + CHECK_FALSE(s.contains(4)); + CHECK_FALSE(s.contains(5)); + CHECK_FALSE(s.contains(7)); + CHECK_FALSE(s.contains(9)); + CHECK_FALSE(r.contains(10)); +} + +TEST_CASE("Range::contains(Range)") { + Range r(5, 10); + + CHECK(r.contains(Range(5, 10))); + CHECK(r.contains(Range(6, 9))); + CHECK(r.contains(Range(7, 7))); // empty range is always contained + CHECK_FALSE(r.contains(Range(4, 10))); + CHECK_FALSE(r.contains(Range(5, 11))); + CHECK_FALSE(r.contains(Range(0, 20))); +} + +TEST_CASE("Range::overlaps") { + Range r(5, 10); + + CHECK(r.overlaps(Range(5, 10))); + CHECK(r.overlaps(Range(0, 6))); + CHECK(r.overlaps(Range(9, 20))); + CHECK_FALSE(r.overlaps(Range(0, 5))); + CHECK_FALSE(r.overlaps(Range(10, 20))); + CHECK_FALSE(r.overlaps(Range(5, 5))); // empty range never overlaps +} + +TEST_CASE("Range::intersection") { + CHECK(Range(0, 10).intersection(Range(5, 15)) == Range(5, 10)); + CHECK(Range(0, 5).intersection(Range(10, 15)).is_empty()); + CHECK(Range(0, 10).intersection(Range(2, 8)) == Range(2, 8)); +} + +TEST_CASE("Range::split") { + SECTION("expected") { + Range r(0, 10); + auto [a, b] = r.split(4); + CHECK(a == Range(0, 4)); + CHECK(b == Range(4, 10)); + } + + SECTION("mid < start") { + Range r(5, 10); + auto [a, b] = r.split(0); + CHECK(a == Range(5, 5)); + CHECK(b == Range(5, 10)); + } + + SECTION("mid > stop") { + Range r(5, 10); + auto [a, b] = r.split(20); + CHECK(a == Range(5, 10)); + CHECK(b == Range(10, 10)); + } +} + +TEST_CASE("Range::operator==/!=") { + CHECK(Range(2, 5) == Range(2, 5)); + CHECK(Range(2, 5) != Range(2, 6)); + CHECK(Range(2, 5) != Range(3, 5)); +} + +TEST_CASE("Range shift") { + CHECK(Range(2, 5) + 3 == Range(5, 8)); + CHECK(3 + Range(2, 5) == Range(5, 8)); + CHECK(Range(2, 5) - 1 == Range(1, 4)); +} + +TEST_CASE("Range::from and converting constructor") { + SECTION("wider") { + Range src(2, 5); + Range dst = Range::from(src); + CHECK(dst == Range(2, 5)); + + Range dst2(src); + CHECK(dst2 == Range(2, 5)); + } + + SECTION("narrowing ok") { + Range src(2, 5); + Range dst(src); + CHECK(dst == Range(2, 5)); + } + + SECTION("narrowing fails") { + Range src(0, 100000); + CHECK_THROWS(Range(src)); + } +} + +TEST_CASE("Range::is_convertible_to") { + CHECK(Range(2, 5).is_convertible_to()); + CHECK(Range(2, 5).is_convertible_to()); + CHECK_FALSE(Range(0, 100000).is_convertible_to()); +} diff --git a/test/core/test_shape.cpp b/test/core/test_shape.cpp new file mode 100644 index 00000000..976808c8 --- /dev/null +++ b/test/core/test_shape.cpp @@ -0,0 +1,172 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/shape.hpp" + +using namespace kmm; + +template +void foo(DomainT x) {} + +#include "kmm/core/bounds.hpp" + +TEST_CASE("Shape construction") { + SECTION("default constructs all zeros") { + Shape<3, int> s; + CHECK(s[0] == 0); + CHECK(s[1] == 0); + CHECK(s[2] == 0); + } + + SECTION("variadic constructor takes exactly N extents") { + Shape<3, int> s(4, 5, 6); + CHECK(s[0] == 4); + CHECK(s[1] == 5); + CHECK(s[2] == 6); + } + + SECTION("wraps an existing Vec unchanged") { + Vec v {4, 5}; + Shape<2, int> s(v); + CHECK(s[0] == 4); + CHECK(s[1] == 5); + } + + SECTION("shape() free function") { + auto s = shape(4, 5, 6); + CHECK(s == Shape<3, int>(4, 5, 6)); + } +} + +TEST_CASE("Shape::fill/one/zero") { + CHECK(Shape<3, int>::fill(7) == Shape<3, int>(7, 7, 7)); + CHECK(Shape<3, int>::one() == Shape<3, int>(1, 1, 1)); + CHECK(Shape<3, int>::zero() == Shape<3, int>(0, 0, 0)); +} + +TEST_CASE("Shape::from") { + SECTION("same dimensionality copies values") { + Vec v {4, 5, 6}; + CHECK(Shape<3, int>::from(v) == Shape<3, int>(4, 5, 6)); + } + + SECTION("growing dimensionality pads with one") { + Vec v {4, 5}; + CHECK(Shape<4, int>::from(v) == Shape<4, int>(4, 5, 1, 1)); + } + + SECTION("shrinking dimensionality truncates") { + Vec v {4, 5, 6, 7}; + CHECK(Shape<2, int>::from(v) == Shape<2, int>(4, 5)); + } +} + +TEST_CASE("Shape converting constructor") { + SECTION("growing dimensionality pads with one") { + Shape<2, int> src(4, 5); + Shape<4, int> dst(src); + CHECK(dst == Shape<4, int>(4, 5, 1, 1)); + } + + SECTION("shrinking dimensionality succeeds when dropped axes are one") { + Shape<3, int> src(4, 5, 1); + Shape<2, int> dst(src); + CHECK(dst == Shape<2, int>(4, 5)); + } + + SECTION("shrinking dimensionality throws when a dropped axis is not one") { + Shape<3, int> src(4, 5, 2); + CHECK_THROWS_AS((Shape<2, int>(src)), std::overflow_error); + } + + SECTION("narrowing type that overflows throws") { + Shape<2, int> src(300, 1); + CHECK_THROWS_AS((Shape<2, signed char>(src)), std::overflow_error); + } + + SECTION("narrowing type that fits succeeds") { + Shape<2, int> src(4, 5); + Shape<2, signed char> dst(src); + CHECK(dst[0] == 4); + CHECK(dst[1] == 5); + } +} + +TEST_CASE("Shape::is_convertible_to") { + CHECK(Shape<2, int>(4, 5).is_convertible_to<2, long>()); + CHECK_FALSE(Shape<2, int>(300, 1).is_convertible_to<2, signed char>()); + CHECK(Shape<3, int>(4, 5, 1).is_convertible_to<2, int>()); + CHECK_FALSE(Shape<3, int>(4, 5, 2).is_convertible_to<2, int>()); +} + +TEST_CASE("Shape::get_or_default") { + Shape<2, int> s(4, 5); + + CHECK(s.get_or_default(0) == 4); + CHECK(s.get_or_default(1) == 5); + CHECK(s.get_or_default(2) == 1); + CHECK(s.get_or_default(2, 99) == 99); + + Shape<0, int> empty; + CHECK(empty.get_or_default(0) == 1); + CHECK(empty.get_or_default(0, 42) == 42); +} + +TEST_CASE("Shape::is_empty") { + CHECK_FALSE(Shape<2, int>(4, 5).is_empty()); + CHECK(Shape<2, int>(0, 5).is_empty()); + CHECK(Shape<2, int>(4, 0).is_empty()); + CHECK(Shape<2, int>(-1, 5).is_empty()); + CHECK_FALSE(Shape<0, int>().is_empty()); +} + +TEST_CASE("Shape::volume") { + CHECK(Shape<3, int>(2, 3, 4).volume() == 24); + CHECK(Shape<2, int>(0, 5).volume() == 0); + CHECK(Shape<0, int>().volume() == 1); + CHECK(Shape<1, int>(7).volume() == 7); +} + +TEST_CASE("Shape::contains") { + Shape<2, int> s(3, 4); + + SECTION("matching dimensionality") { + CHECK(s.contains(Point<2, int>(0, 0))); + CHECK(s.contains(Point<2, int>(2, 3))); + CHECK_FALSE(s.contains(Point<2, int>(3, 0))); + CHECK_FALSE(s.contains(Point<2, int>(0, 4))); + CHECK_FALSE(s.contains(Point<2, int>(-1, 0))); + } + + SECTION("point has more dims: extra dims must be zero") { + CHECK(s.contains(Point<3, int>(1, 2, 0))); + CHECK_FALSE(s.contains(Point<3, int>(1, 2, 1))); + } + + SECTION("shape has more dims: missing point dims default to index 0") { + Shape<3, int> s3(3, 4, 5); + CHECK(s3.contains(Point<2, int>(1, 2))); + + Shape<3, int> s3_empty(3, 4, 0); + CHECK_FALSE(s3_empty.contains(Point<2, int>(1, 2))); + } +} + +TEST_CASE("Shape equality") { + CHECK(Shape<2, int>(4, 5) == Shape<2, int>(4, 5)); + CHECK(Shape<2, int>(4, 5) != Shape<2, int>(4, 6)); + + SECTION("shapes of different dimensionality compare via one-padding") { + CHECK(Shape<2, int>(4, 5) == Shape<3, int>(4, 5, 1)); + CHECK(Shape<2, int>(4, 5) != Shape<3, int>(4, 5, 2)); + } +} + +TEST_CASE("concat(Shape, Shape)") { + Shape<2, int> a(2, 3); + Shape<3, int> b(4, 5, 6); + + Shape<5, int> c = concat(a, b); + CHECK(c == Shape<5, int>(2, 3, 4, 5, 6)); +} diff --git a/test/core/test_strides.cpp b/test/core/test_strides.cpp new file mode 100644 index 00000000..6aedbf0f --- /dev/null +++ b/test/core/test_strides.cpp @@ -0,0 +1,309 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/strides.hpp" + +#define CHECK_TYPE(V, ...) CHECK(std::is_same::value) + +using namespace kmm; + +TEST_CASE("Strides N=0") { + Strides<> s = {}; + + CHECK(Strides<>::rank == 0); + CHECK(s[0] == 0); + CHECK(s.get(ConstIndex<0>()) == 0); + CHECK(s.linearize_offset({}) == 0); + CHECK(fmt::to_string(s) == "{}"); +} + +TEST_CASE("Strides N=1") { + Strides s = {5}; + + CHECK(decltype(s)::rank == 1); + CHECK(s[0] == 5); + CHECK(s[1] == 0); + CHECK(s.get(ConstIndex<0>()) == 5); + CHECK(s.get(ConstIndex<1>()) == 0); + CHECK(s.linearize_offset({2}) == 2 * 5); + CHECK(fmt::to_string(s) == "{5}"); + + Strides> t = {{}}; + + CHECK(decltype(t)::rank == 1); + CHECK(t[0] == 5); + CHECK(t[1] == 0); + CHECK(t.get(ConstIndex<0>()) == 5); + CHECK(t.get(ConstIndex<1>()) == 0); + CHECK(t.linearize_offset({2}) == 2 * 5); + CHECK(fmt::to_string(t) == "{5}"); +} + +TEST_CASE("Strides N=2") { + Strides s = {5, 1}; + + CHECK(decltype(s)::rank == 2); + CHECK(s[0] == 5); + CHECK(s[1] == 1); + CHECK(s[2] == 0); + CHECK(s.get(ConstIndex<0>()) == 5); + CHECK(s.get(ConstIndex<1>()) == 1); + CHECK(s.get(ConstIndex<2>()) == 0); + CHECK(s.linearize_offset({2, 10}) == 5 * 2 + 10); + CHECK(fmt::to_string(s) == "{5, 1}"); + + Strides> t = {5, {}}; + + CHECK(decltype(t)::rank == 2); + CHECK(t[0] == 5); + CHECK(t[1] == 1); + CHECK(t[2] == 0); + CHECK(t.get(ConstIndex<0>()) == 5); + CHECK(t.get(ConstIndex<1>()) == 1); + CHECK(t.get(ConstIndex<2>()) == 0); + CHECK(t.linearize_offset({2, 10}) == 5 * 2 + 10); + CHECK(fmt::to_string(t) == "{5, 1}"); + + Strides, ConstValue<1>> p; + + CHECK(decltype(p)::rank == 2); + CHECK(p[0] == 5); + CHECK(p[1] == 1); + CHECK(p[2] == 0); + CHECK(p.get(ConstIndex<0>()) == 5); + CHECK(p.get(ConstIndex<1>()) == 1); + CHECK(p.get(ConstIndex<2>()) == 0); + CHECK(p.linearize_offset({2, 10}) == 5 * 2 + 10); + CHECK(fmt::to_string(p) == "{5, 1}"); +} + +TEST_CASE("Strides N=3") { + Strides s = {10, 5, 1}; + + CHECK(decltype(s)::rank == 3); + CHECK(s[0] == 10); + CHECK(s[1] == 5); + CHECK(s[2] == 1); + CHECK(s[3] == 0); + CHECK(s.get(ConstIndex<0>()) == 10); + CHECK(s.get(ConstIndex<1>()) == 5); + CHECK(s.get(ConstIndex<2>()) == 1); + CHECK(s.get(ConstIndex<3>()) == 0); + CHECK(s.linearize_offset({6, 2, 10}) == 6 * 10 + 5 * 2 + 10); + CHECK(fmt::to_string(s) == "{10, 5, 1}"); +} + +TEST_CASE("Strides sizeof") { + // N = 0 is empty + CHECK(std::is_empty_v>); + CHECK(sizeof(Strides<>) == 1); + + // Non-empty means N strides + CHECK(sizeof(Strides) == sizeof(long)); + CHECK(sizeof(Strides) == 2 * sizeof(long)); + CHECK(sizeof(Strides) == 3 * sizeof(long)); + + // Only static stride is always empty + CHECK(sizeof(Strides>) == 1); + CHECK(sizeof(Strides, ConstValue<2>>) == 1); + CHECK(sizeof(Strides, ConstValue<2>, ConstValue<3>>) == 1); +} + +TEST_CASE("Strides default constructor") { + SECTION("N=1") { + Strides s; + CHECK(decltype(s)::rank == 1); + CHECK(s[0] == 0); + } + + SECTION("N=2") { + Strides s; + CHECK(s[0] == 0); + CHECK(s[1] == 0); + } + + SECTION("N=2, static values") { + Strides> s; + CHECK(s[0] == 0); + CHECK(s[1] == 5); + } +} + +TEST_CASE("Strides converting constructor") { + SECTION("dynamic to static, success") { + Strides src = {5}; + Strides> dst = src; + CHECK(dst[0] == 5); + } + + SECTION("dynamic to static, failure") { + Strides src = {7}; + CHECK_THROWS_AS((Strides>(src)), std::overflow_error); + } + + SECTION("static to dynamic, success") { + Strides> src; + Strides dst = src; + CHECK(dst[0] == 5); + } + + SECTION("widening, success") { + Strides src = {7}; + Strides dst = src; + CHECK(dst[0] == 7); + } + + SECTION("narrowing, success") { + Strides src = {100}; + Strides dst = src; + CHECK(dst[0] == 100); + } + + SECTION("narrowing, fails") { + Strides src = {300}; + CHECK_THROWS_AS((Strides(src)), std::overflow_error); + } +} + +TEST_CASE("Strides::to_vec") { + SECTION("rank 0") { + Strides<> s; + auto v = s.to_vec(); + CHECK_TYPE(v, Vec); + } + + SECTION("rank 2 with mixed dynamic/static axes") { + Strides> s = {5, {}}; + auto v = s.to_vec(); + + CHECK_TYPE(v, Vec); + CHECK(v[0] == 5); + CHECK(v[1] == 1); + } +} + +TEST_CASE("Strides::operator==/!=") { + Strides<> a; + Strides b = {5}; + Strides c = {0, 0}; + Strides d = {5, 0}; + Strides, ConstValue<1>> e; + + CHECK(a == a); + CHECK(a != b); + CHECK(a == c); + CHECK(a != d); + CHECK(a != e); + + CHECK(b != a); + CHECK(b == b); + CHECK(b != c); + CHECK(b == d); + CHECK(b != e); + + CHECK(c == a); + CHECK(c != b); + CHECK(c == c); + CHECK(c != d); + CHECK(c != e); + + CHECK(d != a); + CHECK(d == b); + CHECK(d != c); + CHECK(d == d); + CHECK(d != e); + + CHECK(e != a); + CHECK(e != b); + CHECK(e != c); + CHECK(e != d); + CHECK(e == e); + + // different data types + Strides f = {5}; + Strides g = {5}; + Strides h = {LONG_MAX}; + + CHECK(f == g); + CHECK(g == f); + CHECK(f != h); + CHECK(h != f); +} + +TEST_CASE("make_strides") { + auto x = make_strides(5, 2, ConstValue<1>()); + + CHECK_TYPE(x, Strides>); + CHECK(x == Strides(5, 2, 1)); + + auto y = make_strides(ConstValue<5>(), ConstValue<1>()); + + CHECK_TYPE(y, Strides, ConstValue<1>>); + CHECK(y == Strides(5, 1)); +} + +TEST_CASE("make_strides_from_shape") { + using index_t = default_index_type; + + SECTION("N=0") { + auto s0 = make_strides_from_shape(Shape()); + CHECK_TYPE(s0, Strides<>); + + auto s0_col = make_strides_from_shape(Shape()); + CHECK_TYPE(s0_col, Strides<>); + } + + SECTION("N=1") { + auto s1 = make_strides_from_shape(Shape(index_t(7))); + CHECK_TYPE(s1, Strides>); + CHECK(s1[0] == 1); + + auto s1_col = make_strides_from_shape(Shape(index_t(7))); + CHECK_TYPE(s1_col, Strides>); + CHECK(s1_col[0] == 1); + } + + SECTION("N=2") { + auto s2 = make_strides_from_shape(Shape(10, 5)); + CHECK_TYPE(s2, Strides>); + CHECK(s2[0] == 5); + CHECK(s2[1] == 1); + + auto s2_col = make_strides_from_shape(Shape(10, 5)); + CHECK_TYPE(s2_col, Strides, index_t>); + CHECK(s2_col[0] == 1); + CHECK(s2_col[1] == 10); + } + + SECTION("N=3") { + auto s3 = make_strides_from_shape(Shape(10, 5, 3)); + CHECK_TYPE(s3, Strides>); + CHECK(s3[0] == 15); + CHECK(s3[1] == 3); + CHECK(s3[2] == 1); + + auto s3_col = make_strides_from_shape(Shape(10, 5, 3)); + CHECK_TYPE(s3_col, Strides, index_t, index_t>); + CHECK(s3_col[0] == 1); + CHECK(s3_col[1] == 10); + CHECK(s3_col[2] == 50); + } + + SECTION("N=3, aligned") { + index_t alignment = 4; + + auto s3 = make_strides_from_shape(Shape(10, 5, 3), alignment); + CHECK_TYPE(s3, Strides>); + CHECK(s3[0] == 20); + CHECK(s3[1] == 4); + CHECK(s3[2] == 1); + + auto s3_col = make_strides_from_shape(Shape(10, 5, 3), alignment); + CHECK_TYPE(s3_col, Strides, index_t, index_t>); + CHECK(s3_col[0] == 1); + CHECK(s3_col[1] == 12); + CHECK(s3_col[2] == 60); + } +} \ No newline at end of file diff --git a/test/core/test_type_utils.cpp b/test/core/test_type_utils.cpp new file mode 100644 index 00000000..cd5c6acd --- /dev/null +++ b/test/core/test_type_utils.cpp @@ -0,0 +1,111 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/core/type_utils.hpp" + +#define CHECK_TYPE(V, ...) CHECK(std::is_same::value) + +using namespace kmm; + +TEST_CASE("IndexSequence methods") { + CHECK(IndexSequence<0, 1, 2>::all([](auto i) { return i() < 3; })); + CHECK_FALSE(IndexSequence<0, 1, 2>::all([](auto i) { return i() < 2; })); + + int sum = 0; + IndexSequence<0, 1, 2, 3>::for_each([&](auto i) { sum += static_cast(i()); }); + CHECK(sum == 6); + + auto v = IndexSequence<0, 1, 2>::construct>([](auto i) { + return static_cast(i()) * 10; + }); + + CHECK(v[0] == 0); + CHECK(v[1] == 10); + CHECK(v[2] == 20); + + auto va = IndexSequence<0, 1, 2>::fill>(7); + CHECK(va[0] == 7); + CHECK(va[1] == 7); + CHECK(va[2] == 7); + + // using Seq = IndexSequence<5, 2, 8>; + // CHECK(Seq::contains<5>); + // CHECK(Seq::contains<2>); + // CHECK(Seq::contains<8>); + // CHECK_FALSE(Seq::contains<0>); + // + // CHECK(Seq::location_of<5> == 0); + // CHECK(Seq::location_of<2> == 1); + // CHECK(Seq::location_of<8> == 2); +} + +TEST_CASE("make_index_sequence") { + CHECK_TYPE(make_index_sequence<0>(), IndexSequence<>); + CHECK_TYPE(make_index_sequence<1>(), IndexSequence<0>); + CHECK_TYPE(make_index_sequence<2>(), IndexSequence<0, 1>); + CHECK_TYPE(make_index_sequence<3>(), IndexSequence<0, 1, 2>); + CHECK_TYPE(make_index_sequence<4>(), IndexSequence<0, 1, 2, 3>); +} + +TEST_CASE("range_index_sequence_t") { + CHECK_TYPE((range_index_sequence_t<2, 5>()), IndexSequence<2, 3, 4>); + CHECK_TYPE((range_index_sequence_t<4, 4>()), IndexSequence<>); +} + +TEST_CASE("reverse_index_sequence") { + CHECK_TYPE((reverse_index_sequence<4>()), IndexSequence<3, 2, 1, 0>); +} + +TEST_CASE("drop_index_sequence") { + CHECK_TYPE((drop_index_sequence<4, 0>()), IndexSequence<1, 2, 3>); + CHECK_TYPE((drop_index_sequence<4, 1>()), IndexSequence<0, 2, 3>); + CHECK_TYPE((drop_index_sequence<4, 2>()), IndexSequence<0, 1, 3>); + CHECK_TYPE((drop_index_sequence<4, 3>()), IndexSequence<0, 1, 2>); +} + +TEST_CASE("move_axis_to_position_index_sequence") { + // identity: axis already at its target position + CHECK_TYPE((move_axis_to_position_index_sequence<4, 2, 2>()), IndexSequence<0, 1, 2, 3>); + + // move axis 1 to position 3 (rightward): axes 2, 3 shift left to fill the gap + CHECK_TYPE((move_axis_to_position_index_sequence<4, 1, 3>()), IndexSequence<0, 2, 3, 1>); + + // move axis 3 to position 1 (leftward): axes 1, 2 shift right to fill the gap + CHECK_TYPE((move_axis_to_position_index_sequence<4, 3, 1>()), IndexSequence<0, 3, 1, 2>); + + // move to front / back + CHECK_TYPE((move_axis_to_position_index_sequence<3, 2, 0>()), IndexSequence<2, 0, 1>); + CHECK_TYPE((move_axis_to_position_index_sequence<3, 0, 2>()), IndexSequence<1, 2, 0>); + + // single-axis and no-op cases + CHECK_TYPE((move_axis_to_position_index_sequence<1, 0, 0>()), IndexSequence<0>); + CHECK_TYPE((move_axis_to_position_index_sequence<4, 0, 0>()), IndexSequence<0, 1, 2, 3>); + CHECK_TYPE((move_axis_to_position_index_sequence<4, 3, 3>()), IndexSequence<0, 1, 2, 3>); +} + +TEST_CASE("is_partial_permutation") { + CHECK((is_partial_permutation, 0>)); + CHECK((is_partial_permutation, 1>)); + CHECK((is_partial_permutation, 3>)); + CHECK((is_partial_permutation, 3>)); + CHECK((is_partial_permutation, 4>)); // partial: doesn't cover 0 or 2 + + CHECK_FALSE((is_partial_permutation, 1>)); + CHECK_FALSE((is_partial_permutation, 3>)); + CHECK_FALSE((is_partial_permutation, 3>)); + + // out of bounds for the given N, even though otherwise injective + CHECK_FALSE((is_partial_permutation, 3>)); +} + +TEST_CASE("is_permutation") { + CHECK(is_permutation>); + CHECK(is_permutation>); + CHECK(is_permutation>); + CHECK(is_permutation>); + + CHECK_FALSE(is_permutation>); // repeated index + CHECK_FALSE(is_permutation>); // injective, but doesn't cover 0..size-1 +} diff --git a/test/core/test_vec.cpp b/test/core/test_vec.cpp new file mode 100644 index 00000000..1c069683 --- /dev/null +++ b/test/core/test_vec.cpp @@ -0,0 +1,137 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/core/vec.hpp" + +using namespace kmm; + +TEST_CASE("Vec") { + Vec a; + Vec b {}; + auto c = Vec {}; + + // CHECK(a.data() == nullptr); + // CHECK(b.data() == nullptr); + // CHECK(c.data() == nullptr); +} + +TEST_CASE("Vec") { + Vec a; + Vec b {1}; + + CHECK(b[0] == 1); + CHECK(b.x == 1); + // CHECK(b.data() == &b.x); + + b[0] = 5; + CHECK(b.x == 5); +} + +TEST_CASE("Vec") { + Vec a; + Vec b {1, 2}; + + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b.x == 1); + CHECK(b.y == 2); + + // CHECK(b.data() == &b.x); + // CHECK(b.data()[1] == b.y); + + b[1] = 9; + CHECK(b.y == 9); +} + +TEST_CASE("Vec") { + Vec a; + Vec b {1, 2, 3}; + + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); + CHECK(b.x == 1); + CHECK(b.y == 2); + CHECK(b.z == 3); + + // CHECK(b.data() == &b.x); + // CHECK(b.data()[2] == b.z); + + b[2] = 42; + CHECK(b.z == 42); +} + +TEST_CASE("Vec") { + Vec a; + Vec b {1, 2, 3, 4}; + + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); + CHECK(b[3] == 4); + CHECK(b.x == 1); + CHECK(b.y == 2); + CHECK(b.z == 3); + CHECK(b.w == 4); + + // CHECK(b.data() == &b.x); + // CHECK(b.data()[3] == b.w); + + b[3] = 42; + CHECK(b.w == 42); +} + +TEST_CASE("Vec generic") { + Vec a; + Vec b {1, 2, 3, 4, 5}; + + for (size_t i = 0; i < 5; i++) { + CHECK(b[i] == static_cast(i + 1)); + } + + // CHECK(b.data() == b.values); + + b[4] = 100; + CHECK(b.values[4] == 100); + + const Vec c {5, 4, 3, 2, 1}; + for (size_t i = 0; i < 5; i++) { + CHECK(c[i] == static_cast(5 - i)); + } +} + +TEST_CASE("fill") { + auto v0 = fill<0>(42); + // CHECK(v0.data() == nullptr); + + auto v1 = fill<1>(7); + CHECK(v1[0] == 7); + + auto v2 = fill<2>(3); + CHECK(v2[0] == 3); + CHECK(v2[1] == 3); + + auto v4 = fill<4>(9); + for (size_t i = 0; i < 4; i++) { + CHECK(v4[i] == 9); + } + + auto v5 = fill<6>(1); + for (size_t i = 0; i < 6; i++) { + CHECK(v5[i] == 1); + } +} + +TEST_CASE("concat") { + auto a = Vec {}; + auto b = Vec {1}; + auto c = Vec {2, 3, 4}; + + auto x = concat(a, b); + CHECK(x[0] == 1); + + auto y = concat(b, c); + CHECK(y[0] == 1); + CHECK(y[1] == 2); + CHECK(y[2] == 3); + CHECK(y[3] == 4); +} diff --git a/test/core/test_view.cpp b/test/core/test_view.cpp new file mode 100644 index 00000000..7e319ed5 --- /dev/null +++ b/test/core/test_view.cpp @@ -0,0 +1,272 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/core/view.hpp" + +using namespace kmm; + +TEST_CASE("NDView default construction") { + ViewMut view; + + CHECK(view.data() == nullptr); + CHECK(view.is_empty()); +} + +TEST_CASE("make_view/NDView basics") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + + CHECK(view.shape() == Shape<2, int>(4, 5)); + CHECK(view.extent(0) == 4); + CHECK(view.extent(1) == 5); + CHECK(view.size() == 20); + CHECK_FALSE(view.is_empty()); + CHECK(view.stride(0) == 5); + CHECK(view.stride(1) == 1); + CHECK(view.strides() == Vec(5, 1)); + CHECK(view.data() == data); + CHECK(view.is_contiguous()); + CHECK(view.is_contiguous(MemoryOrder::RowMajor)); + CHECK_FALSE(view.is_contiguous(MemoryOrder::ColMajor)); + + SECTION("an axis of extent zero makes the view empty") { + auto empty_view = make_view(data, Shape<2, int>(0, 5)); + CHECK(empty_view.is_empty()); + CHECK(empty_view.size() == 0); + } +} + +TEST_CASE("make_view with an explicit stride policy") { + int data[6] = {0, 1, 2, 3, 4, 5}; + auto view = make_view(data, Shape<2, int>(2, 3)); + + CHECK(view.stride(0) == 1); + CHECK(view.stride(1) == 2); + CHECK(view.is_contiguous(MemoryOrder::ColMajor)); + CHECK_FALSE(view.is_contiguous(MemoryOrder::RowMajor)); +} + +TEST_CASE("NDView::access") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + + CHECK(view(1, 2) == data[1 * 5 + 2]); + CHECK(view[Vec(1, 2)] == data[1 * 5 + 2]); + CHECK(view[1][2] == data[1 * 5 + 2]); + CHECK(view.access(Vec(1, 2)) == data[1 * 5 + 2]); + + SECTION("writes propagate back to the underlying storage") { + view(0, 0) = 42; + CHECK(data[0] == 42); + } + + SECTION("contains") { + CHECK(view.contains(Vec(0, 0))); + CHECK(view.contains(Vec(3, 4))); + CHECK_FALSE(view.contains(Vec(4, 0))); + CHECK_FALSE(view.contains(Vec(0, 5))); + } +} + +TEST_CASE("NDView::move_origin") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + + SECTION("shifting the first axis") { + auto moved = view.move_origin(Vec(1, 0)); + CHECK(moved.shape() == Shape<2, int>(4, 5)); + + for (int i = 0; i < 3; i++) { + for (int j = 0; j < 5; j++) { + CHECK(moved(i, j) == view(i + 1, j)); + } + } + } + + SECTION("shifting the second axis") { + auto moved = view.move_origin(Vec(0, 1)); + + for (int i = 0; i < 4; i++) { + for (int j = 0; j < 4; j++) { + CHECK(moved(i, j) == view(i, j + 1)); + } + } + } +} + +TEST_CASE("NDView::restrict_bounds/restrict_axis") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + + SECTION("restrict_axis narrows one axis, keeping absolute indices") { + auto restricted = view.restrict_axis<0>(1, 3); + + CHECK(restricted.extent(0) == 2); + CHECK(restricted.extent(1) == 5); + CHECK(restricted(1, 2) == view(1, 2)); + CHECK(restricted(2, 3) == view(2, 3)); + } + + SECTION("restrict_bounds narrows to the intersection with the given bounds") { + auto restricted = view.restrict_bounds(Bounds<2, int>(Range(1, 3), Range(0, 5))); + + CHECK(restricted.extent(0) == 2); + CHECK(restricted.extent(1) == 5); + CHECK(restricted(1, 0) == view(1, 0)); + CHECK(restricted(2, 4) == view(2, 4)); + } +} + +TEST_CASE("NDView::zero_origin") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + auto restricted = view.restrict_axis<0>(1, 3); + auto zo = restricted.zero_origin(); + + CHECK(zo.shape() == Shape<2, int>(2, 5)); + CHECK(zo(0, 0) == view(1, 0)); + CHECK(zo(1, 4) == view(2, 4)); +} + +TEST_CASE("NDView::drop_axis") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + auto dropped = view.drop_axis<1>(2); + + CHECK(dropped.shape() == Shape<1, int>(4)); + + for (int i = 0; i < 4; i++) { + CHECK(dropped(i) == view(i, 2)); + } +} + +TEST_CASE("NDView::insert_axis") { + int data[4] = {10, 20, 30, 40}; + auto view = make_view(data, Shape<1, int>(4)); + + auto inserted = view.insert_axis<1>(3); + CHECK(inserted.shape() == Shape<2, int>(4, 3)); + + for (int i = 0; i < 4; i++) { + for (int j = 0; j < 3; j++) { + CHECK(inserted(i, j) == view(i)); + } + } +} + +TEST_CASE("NDView::reverse_axes") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + auto reversed = view.reverse_axes(); + + CHECK(reversed.shape() == Shape<2, int>(5, 4)); + + for (int i = 0; i < 4; i++) { + for (int j = 0; j < 5; j++) { + CHECK(reversed(j, i) == view(i, j)); + } + } +} + +TEST_CASE("NDView::slice_axis") { + int data[20]; + for (int i = 0; i < 20; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<2, int>(4, 5)); + + SECTION("with a start/end pair narrows the axis and rebases it to zero") { + auto sliced = view.slice_axis<0>(1, 3); + CHECK(sliced.shape() == Shape<2, int>(2, 5)); + + for (int i = 0; i < 2; i++) { + for (int j = 0; j < 5; j++) { + CHECK(sliced(i, j) == view(i + 1, j)); + } + } + } + + SECTION("with a Range token narrows the axis the same way") { + auto sliced = view.slice_axis<1>(Range(2, 4)); + CHECK(sliced.shape() == Shape<2, int>(4, 2)); + + for (int i = 0; i < 4; i++) { + for (int j = 0; j < 2; j++) { + CHECK(sliced(i, j) == view(i, j + 2)); + } + } + } + + SECTION("with all leaves the axis unchanged") { + auto sliced = view.slice_axis<0>(all); + CHECK(sliced.shape() == view.shape()); + } +} + +TEST_CASE("NDView::slice") { + int data[24]; + for (int i = 0; i < 24; i++) { + data[i] = i; + } + + auto view = make_view(data, Shape<3, int>(2, 3, 4)); + + SECTION("mixing a plain index, all, and a Range across axes") { + auto sliced = view.slice(1, all, Range(1, 3)); + CHECK(sliced.shape() == Shape<2, int>(3, 2)); + + CHECK(sliced.layout().base_offset() == view.layout().base_offset() + 1 * 12 + 1); + CHECK(sliced.data() == view.data()); + + for (int j = 0; j < 3; j++) { + for (int k = 0; k < 2; k++) { + CHECK(sliced(j, k) == view(1, j, k + 1)); + } + } + } + + SECTION("an index for every axis drops down to a rank-0 view") { + auto sliced = view.slice(1, 2, 3); + CHECK(decltype(sliced)::rank == 0); + CHECK(sliced.size() == 1); + CHECK(sliced() == view(1, 2, 3)); + } +} + +TEST_CASE("NDView converting constructor") { + int data[6] = {0, 1, 2, 3, 4, 5}; + auto view = make_view(data, Shape<2, int>(2, 3)); + + NDView const_view = view; + + CHECK(const_view.data() == data); + CHECK(const_view(1, 2) == 5); +} diff --git a/test/dag/test_data_distribution.cpp b/test/dag/test_data_distribution.cpp deleted file mode 100644 index 70ed48c2..00000000 --- a/test/dag/test_data_distribution.cpp +++ /dev/null @@ -1,106 +0,0 @@ -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/core/distribution.hpp" - -using namespace kmm; - -TEST_CASE("Distribution<1>") { - std::vector> chunks = { - ArrayChunk<1> {.owner_id = DeviceId(0), .offset = 0, .size = 10}, - ArrayChunk<1> {.owner_id = DeviceId(1), .offset = 10, .size = 10}, - ArrayChunk<1> {.owner_id = DeviceId(2), .offset = 20, .size = 6} - }; - - SECTION("no swap") {} - - SECTION("swap 0 and 1") { - std::swap(chunks[0], chunks[1]); - } - - SECTION("swap 0 and 2") { - std::swap(chunks[0], chunks[2]); - } - - SECTION("swap 1 and 2") { - std::swap(chunks[1], chunks[2]); - } - - auto dist = Distribution<1>::from_chunks(26, chunks); - - CHECK(dist.num_chunks() == 3); - CHECK(dist.chunk_size() == Dim {10}); - CHECK(dist.array_size() == Dim {26}); - - CHECK(dist.chunk(0).offset == 0); - CHECK(dist.chunk(0).size == 10); - CHECK(dist.chunk(0).owner_id == DeviceId(0)); - - CHECK(dist.chunk(1).offset == 10); - CHECK(dist.chunk(1).size == 10); - CHECK(dist.chunk(1).owner_id == DeviceId(1)); - - CHECK(dist.chunk(2).offset == 20); - CHECK(dist.chunk(2).size == 6); - CHECK(dist.chunk(2).owner_id == DeviceId(2)); -} - -TEST_CASE("Distribution<2>") { - std::vector> chunks = { - ArrayChunk<2> {.owner_id = DeviceId(1), .offset = {0, 0}, .size = {15, 10}}, - ArrayChunk<2> {.owner_id = DeviceId(2), .offset = {0, 10}, .size = {15, 10}}, - ArrayChunk<2> {.owner_id = DeviceId(3), .offset = {0, 20}, .size = {15, 7}}, - ArrayChunk<2> {.owner_id = DeviceId(4), .offset = {15, 0}, .size = {14, 10}}, - ArrayChunk<2> {.owner_id = DeviceId(5), .offset = {15, 10}, .size = {14, 10}}, - ArrayChunk<2> {.owner_id = DeviceId(6), .offset = {15, 20}, .size = {14, 7}} - }; - - SECTION("no swap") {} - - SECTION("swap 0 and 1") { - std::swap(chunks[0], chunks[1]); - } - - SECTION("swap 0 and 4") { - std::swap(chunks[0], chunks[4]); - } - - SECTION("swap 3 and 5") { - std::swap(chunks[3], chunks[5]); - } - - SECTION("swap 3 and 4") { - std::swap(chunks[3], chunks[4]); - } - - auto dist = Distribution<2>::from_chunks({29, 27}, chunks); - - CHECK(dist.num_chunks() == 6); - CHECK(dist.chunk_size() == Dim {15, 10}); - CHECK(dist.array_size() == Dim {29, 27}); - - CHECK(dist.chunk(0).offset == Point {0, 0}); - CHECK(dist.chunk(0).size == Dim {15, 10}); - CHECK(dist.chunk(0).owner_id == DeviceId(1)); - - CHECK(dist.chunk(1).offset == Point {0, 10}); - CHECK(dist.chunk(1).size == Dim {15, 10}); - CHECK(dist.chunk(1).owner_id == DeviceId(2)); - - CHECK(dist.chunk(2).offset == Point {0, 20}); - CHECK(dist.chunk(2).size == Dim {15, 7}); - CHECK(dist.chunk(2).owner_id == DeviceId(3)); - - CHECK(dist.chunk(3).offset == Point {15, 0}); - CHECK(dist.chunk(3).size == Dim {14, 10}); - CHECK(dist.chunk(3).owner_id == DeviceId(4)); - - CHECK(dist.chunk(4).offset == Point {15, 10}); - CHECK(dist.chunk(4).size == Dim {14, 10}); - CHECK(dist.chunk(4).owner_id == DeviceId(5)); - - CHECK(dist.chunk(5).offset == Point {15, 20}); - CHECK(dist.chunk(5).size == Dim {14, 7}); - CHECK(dist.chunk(5).owner_id == DeviceId(6)); -} \ No newline at end of file diff --git a/test/memops/test_copy.cpp b/test/memops/test_copy.cpp new file mode 100644 index 00000000..bc5fd703 --- /dev/null +++ b/test/memops/test_copy.cpp @@ -0,0 +1,207 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memops/copy.hpp" + +using namespace kmm; +using namespace kmm::memops; + +TEST_CASE("copy (CPU)") { + std::vector src = {1, 2, 3, 4, 5, 6}; + std::vector dst(6, 0); + + CopyDescription description(sizeof(int)); + description.add_dimension(6, sizeof(int), sizeof(int)); + + copy(src.data(), dst.data(), description); + + CHECK(dst == src); +} + +TEST_CASE("CopyDescription::src_range/dst_range") { + SECTION("no dimensions") { + CopyDescription description(sizeof(int)); + description.src_offset = 4; + description.dst_offset = 8; + + CHECK(description.src_range() == Range(4, 4 + sizeof(int))); + CHECK(description.dst_range() == Range(8, 8 + sizeof(int))); + } + + SECTION("positive strides") { + CopyDescription description(sizeof(int)); + description.add_dimension(4, sizeof(int), 2 * sizeof(int)); + + CHECK(description.src_range() == Range(0, 4 * sizeof(int))); + CHECK(description.dst_range() == Range(0, 3 * 2 * sizeof(int) + sizeof(int))); + } + + SECTION("negative stride") { + CopyDescription description(sizeof(int)); + description.add_dimension(4, -static_cast(sizeof(int)), sizeof(int)); + + CHECK( + description.src_range() + == Range(-3 * static_cast(sizeof(int)), sizeof(int)) + ); + CHECK(description.dst_range() == Range(0, 4 * sizeof(int))); + } +} + +namespace { +// Minimal duck-typed stand-in for `kmm::Layout`, avoiding a dependency on `kmm/core/layout.hpp`. +template +struct FakeLayout { + static constexpr size_t rank = N; + + ptrdiff_t offset = 0; + ptrdiff_t extents[N]; + ptrdiff_t strides[N]; + ptrdiff_t origins[N] = {}; + + ptrdiff_t base_offset() const { + return offset; + } + + ptrdiff_t extent(size_t axis) const { + return extents[axis]; + } + + ptrdiff_t stride(size_t axis) const { + return strides[axis]; + } + + ptrdiff_t begin(size_t axis) const { + return origins[axis]; + } +}; +} // namespace + +TEST_CASE("make_copy_description") { + // Both `dst` and `src` are contiguous here, so `add_dimension` folds the two axes into one. + FakeLayout<2> dst {/* offset */ 3, /* extents */ {2, 3}, /* strides */ {3, 1}}; + FakeLayout<2> src {/* offset */ 5, /* extents */ {2, 3}, /* strides */ {6, 2}}; + + CopyDescription description = make_copy_description(dst, src, sizeof(int)); + + CHECK(description.element_size == sizeof(int)); + CHECK(description.dst_offset == 3 * static_cast(sizeof(int))); + CHECK(description.src_offset == 5 * static_cast(sizeof(int))); + CHECK(description.num_dims == 1); + + CHECK(description.dims[0].extent == 6); + CHECK(description.dims[0].dst_stride == 1 * static_cast(sizeof(int))); + CHECK(description.dims[0].src_stride == 2 * static_cast(sizeof(int))); +} + +TEST_CASE("make_copy_description (non-contiguous)") { + // `dst` is not contiguous (there is a gap of 1 element between rows), so no merge happens. + FakeLayout<2> dst {/* offset */ 3, /* extents */ {2, 3}, /* strides */ {4, 1}}; + FakeLayout<2> src {/* offset */ 5, /* extents */ {2, 3}, /* strides */ {6, 2}}; + + CopyDescription description = make_copy_description(dst, src, sizeof(int)); + + CHECK(description.num_dims == 2); + + CHECK(description.dims[0].extent == 2); + CHECK(description.dims[0].dst_stride == 4 * static_cast(sizeof(int))); + CHECK(description.dims[0].src_stride == 6 * static_cast(sizeof(int))); + + CHECK(description.dims[1].extent == 3); + CHECK(description.dims[1].dst_stride == 1 * static_cast(sizeof(int))); + CHECK(description.dims[1].src_stride == 2 * static_cast(sizeof(int))); +} + +TEST_CASE("CopyDescription::simplify sorts axes without merging") { + // Axes are added out of order but are not contiguous with each other, so `simplify` must + // sort them into descending stride order without dropping any of them. + CopyDescription description(sizeof(int)); + description.add_dimension(4, 7 * sizeof(int), 7 * sizeof(int)); + description.add_dimension(2, 101 * sizeof(int), 101 * sizeof(int)); + description.add_dimension(3, 23 * sizeof(int), 23 * sizeof(int)); + + CopyDescription result = description.simplify(); + + REQUIRE(result.num_dims == 3); + CHECK(result.dims[0].extent == 2); + CHECK(result.dims[1].extent == 3); + CHECK(result.dims[2].extent == 4); +} + +TEST_CASE("CopyDescription::simplify sorts and merges contiguous axes") { + // Axes are added out of order, and two of them (stride 6 and stride 2) are contiguous with + // each other. `simplify` must sort them into descending stride order first, then merge the + // two contiguous ones, leaving the unrelated third axis (stride 100) untouched. + CopyDescription description(sizeof(int)); + description.add_dimension(3, 2 * sizeof(int), 2 * sizeof(int)); + description.add_dimension(2, 100 * sizeof(int), 100 * sizeof(int)); + description.add_dimension(4, 6 * sizeof(int), 6 * sizeof(int)); + + CopyDescription result = description.simplify(); + + REQUIRE(result.num_dims == 2); + + CHECK(result.dims[0].extent == 2); + CHECK(result.dims[0].src_stride == 100 * static_cast(sizeof(int))); + CHECK(result.dims[0].dst_stride == 100 * static_cast(sizeof(int))); + + CHECK(result.dims[1].extent == 12); + CHECK(result.dims[1].src_stride == 2 * static_cast(sizeof(int))); + CHECK(result.dims[1].dst_stride == 2 * static_cast(sizeof(int))); +} + +TEST_CASE("CopyDescription::simplify normalizes negative strides") { + // An axis walked backwards on both sides. `simplify` flips its iteration order (shifting the + // offsets to the last element and negating the strides), after which the axis is contiguous + // with the element and gets folded away entirely. + CopyDescription description(sizeof(int)); + description.add_dimension( + 4, + -static_cast(sizeof(int)), + -static_cast(sizeof(int)) + ); + + CopyDescription result = description.simplify(); + + CHECK(result.num_dims == 0); + CHECK(result.element_size == 4 * sizeof(int)); + CHECK(result.src_offset == -3 * static_cast(sizeof(int))); + CHECK(result.dst_offset == -3 * static_cast(sizeof(int))); +} + +TEST_CASE("CopyDescription::simplify normalizes a negative stride on one side only") { + // The axis is forward in the source but reversed in the destination. `simplify` keys on the + // (non-zero) `dst_stride`, so it flips the iteration order to make `dst_stride` positive, + // leaving `src_stride` negative, and still merges the two contiguous axes. + CopyDescription description(sizeof(int)); + description.add_dimension( + 3, + 4 * sizeof(int), + -4 * static_cast(sizeof(int)) + ); + description.add_dimension( + 4, + 1 * sizeof(int), + -1 * static_cast(sizeof(int)) + ); + + CopyDescription result = description.simplify(); + + REQUIRE(result.num_dims == 1); + CHECK(result.dims[0].extent == 12); + CHECK(result.dims[0].src_stride == -1 * static_cast(sizeof(int))); + CHECK(result.dims[0].dst_stride == 1 * static_cast(sizeof(int))); +} + +TEST_CASE("CopyDescription::simplify drops axes with extent one") { + CopyDescription description(sizeof(int)); + description.add_dimension(4, 5 * sizeof(int), 5 * sizeof(int)); + description.add_dimension(1, 999 * sizeof(int), 999 * sizeof(int)); + + CopyDescription result = description.simplify(); + + REQUIRE(result.num_dims == 1); + CHECK(result.dims[0].extent == 4); + CHECK(result.dims[0].src_stride == 5 * static_cast(sizeof(int))); +} diff --git a/test/memops/test_fill.cpp b/test/memops/test_fill.cpp new file mode 100644 index 00000000..edf020ea --- /dev/null +++ b/test/memops/test_fill.cpp @@ -0,0 +1,164 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memops/fill.hpp" + +using namespace kmm; +using namespace kmm::memops; + +static void test_fill(const FillDescription& desc, std::byte* data, size_t axis = 0) { + if (axis >= desc.num_dims) { + std::copy_n(desc.value.buffer, desc.value.length, &data[desc.offset]); + } else { + for (memops_extent_type i = 0; i < desc.dims[axis].extent; i++) { + test_fill(desc, data + desc.dims[axis].stride * i, axis + 1); + } + } +} + +static size_t size_fill(const FillDescription& desc, size_t offset = 0, size_t axis = 0) { + if (axis >= desc.num_dims) { + offset += desc.offset + desc.value.length; + } else { + for (memops_extent_type i = 0; i < desc.dims[axis].extent; i++) { + offset = + std::max(offset, size_fill(desc, offset + desc.dims[axis].stride * i, axis + 1)); + } + } + + return offset; +} + +static bool check_fill(const FillDescription& desc) { + size_t size = size_fill(desc) * 2; // *2 just to be sure + + std::vector expected(size); + std::vector actual(size); + std::vector simplified(size); + + test_fill(desc, expected.data()); + test_fill(desc.simplify(), simplified.data()); + memops::fill(actual.data(), desc); + + return actual == expected && simplified == expected; +} + +TEST_CASE("memops::fill matches reference") { + FillValue value = FillValue::from(0x12345678); + + SECTION("scalar, no dimensions") { + FillDescription desc(value); + CHECK(check_fill(desc)); + } + + SECTION("scalar at an offset") { + FillDescription desc(value); + desc.offset = 40; + CHECK(check_fill(desc)); + } + + SECTION("1D contiguous") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/16, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("1D mis-aligned") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/16, /*stride=*/5); + CHECK(check_fill(desc)); + } + + SECTION("2D strided with offset") { + FillDescription desc(value); + desc.offset = 8; + desc.add_dimension(/*extent=*/6, /*stride=*/64); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("3D") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/40); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("4D") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/2, /*stride=*/1500); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/40); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("axes contiguous (simplify merges them)") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/8, /*stride=*/16); + desc.add_dimension(/*extent=*/4, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("axes contiguous (reversed order)") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/4, /*stride=*/4); + desc.add_dimension(/*extent=*/8, /*stride=*/16); + CHECK(check_fill(desc)); + } + + SECTION("axes overlapping (simplify merges them)") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/8, /*stride=*/16); + desc.add_dimension(/*extent=*/4, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("axis with extent one") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/1, /*stride=*/999); + desc.add_dimension(/*extent=*/10, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("negative stride") { + FillDescription desc(value); + desc.offset = 4 * 7; + desc.add_dimension(/*extent=*/8, /*stride=*/-4); + CHECK(check_fill(desc)); + } + + SECTION("2D negative stride") { + FillDescription desc(value); + desc.offset = 256; + desc.add_dimension(/*extent=*/8, /*stride=*/-4); + desc.add_dimension(/*extent=*/2, /*stride=*/-40); + CHECK(check_fill(desc)); + } + + SECTION("zero stride") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/0); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("mixed-sign strides") { + FillDescription desc(value); + desc.offset = 3 * 40; + desc.add_dimension(/*extent=*/4, /*stride=*/-40); + desc.add_dimension(/*extent=*/9, /*stride=*/4); + CHECK(check_fill(desc)); + } + + SECTION("wide fill value") { + FillDescription desc(FillValue::from(3.14159)); + desc.add_dimension(/*extent=*/7, /*stride=*/8); + desc.add_dimension(/*extent=*/2, /*stride=*/64); + CHECK(check_fill(desc)); + } +} diff --git a/test/memops/test_fill_gpu.cu b/test/memops/test_fill_gpu.cu new file mode 100644 index 00000000..0442ac31 --- /dev/null +++ b/test/memops/test_fill_gpu.cu @@ -0,0 +1,163 @@ +#include +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memops/fill.hpp" +#include "kmm/runtime/memops/fill_gpu.hpp" +#include "kmm/utils/gpu_utils.hpp" + +using namespace kmm; +using namespace kmm::memops; + +static void test_fill(const FillDescription& desc, std::byte* data, size_t axis = 0) { + if (axis >= desc.num_dims) { + std::copy_n(desc.value.buffer, desc.value.length, &data[desc.offset]); + } else { + for (memops_extent_type i = 0; i < desc.dims[axis].extent; i++) { + test_fill(desc, data + desc.dims[axis].stride * i, axis + 1); + } + } +} + +static size_t size_fill(const FillDescription& desc, size_t offset = 0, size_t axis = 0) { + if (axis >= desc.num_dims) { + offset += desc.offset + desc.value.length; + } else { + for (memops_extent_type i = 0; i < desc.dims[axis].extent; i++) { + offset = + std::max(offset, size_fill(desc, offset + desc.dims[axis].stride * i, axis + 1)); + } + } + + return offset; +} + +bool check_fill_gpu(const FillDescription& desc) { + g_device_t device = 0; + g_context_t context = nullptr; + g_stream_t stream = nullptr; + + KMM_GPU_CHECK(g_init(0)); + KMM_GPU_CHECK(g_device_get(&device, 0)); + KMM_GPU_CHECK(g_device_primary_ctx_retain(&context, device)); + KMM_GPU_CHECK(g_ctx_push_current(context)); + + size_t size = size_fill(desc) * 2; // *2 just to be sure + + std::vector reference(size); + test_fill(desc, reference.data()); + + g_device_ptr_t dptr = 0; + KMM_GPU_CHECK(g_mem_alloc(&dptr, size)); + KMM_GPU_CHECK(g_memset_d8_async(dptr, 0, size, stream)); + + memops::fill_gpu(stream, reinterpret_cast(dptr), desc); + KMM_GPU_CHECK(g_stream_synchronize(stream)); + + std::vector actual(size); + KMM_GPU_CHECK(g_memcpy_d_to_h(actual.data(), dptr, size)); + + KMM_GPU_CHECK(g_mem_free(dptr)); + KMM_GPU_CHECK(g_ctx_pop_current(&context)); + KMM_GPU_CHECK(g_device_primary_ctx_release(device)); + + return actual == reference; +} + +TEST_CASE("memops::fill_gpu", "[GPU]") { + FillValue value = FillValue::from(0x12345678); + + SECTION("scalar, no dimensions") { + FillDescription desc(value); + CHECK(check_fill_gpu(desc)); + } + + SECTION("scalar at an offset") { + FillDescription desc(value); + desc.offset = 40; + CHECK(check_fill_gpu(desc)); + } + + SECTION("1D contiguous") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/16, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("1D mis-aligned") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/16, /*stride=*/5); + CHECK(check_fill_gpu(desc)); + } + + SECTION("2D strided with offset") { + FillDescription desc(value); + desc.offset = 8; + desc.add_dimension(/*extent=*/6, /*stride=*/64); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("3D") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/40); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("4D") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/2, /*stride=*/1500); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/40); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("axes contiguous (simplify merges them)") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/8, /*stride=*/16); + desc.add_dimension(/*extent=*/4, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("axis with extent one") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/1, /*stride=*/999); + desc.add_dimension(/*extent=*/10, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("negative stride") { + FillDescription desc(value); + desc.offset = 4 * 7; + desc.add_dimension(/*extent=*/8, /*stride=*/-4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("zero stride") { + FillDescription desc(value); + desc.add_dimension(/*extent=*/3, /*stride=*/400); + desc.add_dimension(/*extent=*/4, /*stride=*/0); + desc.add_dimension(/*extent=*/5, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("mixed-sign strides") { + FillDescription desc(value); + desc.offset = 3 * 40; + desc.add_dimension(/*extent=*/4, /*stride=*/-40); + desc.add_dimension(/*extent=*/9, /*stride=*/4); + CHECK(check_fill_gpu(desc)); + } + + SECTION("wide fill value") { + FillDescription desc(FillValue::from(3.14159)); + desc.add_dimension(/*extent=*/7, /*stride=*/8); + desc.add_dimension(/*extent=*/2, /*stride=*/64); + CHECK(check_fill_gpu(desc)); + } +} diff --git a/test/memops/test_reduction.cpp b/test/memops/test_reduction.cpp new file mode 100644 index 00000000..1b0ade9d --- /dev/null +++ b/test/memops/test_reduction.cpp @@ -0,0 +1,146 @@ +#include +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memops/reduction.hpp" + +using namespace kmm; +using namespace kmm::memops; + +TEST_CASE("reduce (CPU)") { + // Reduce a (4, 3) row-major input down to 4 outputs by summing over the trailing axis. + std::vector src = { + 1, + 2, + 3, // + 4, + 5, + 6, // + 7, + 8, + 9, // + 10, + 11, + 12 // + }; + std::vector dst(4, 0.0f); + + ReductionDescription description(DataType::Float32, ReductionOp::Sum); + description.add_dimension(4, 3 * sizeof(float), sizeof(float)); + description.reduction_extent = 3; + description.reduction_stride = sizeof(float); + + reduce(src.data(), dst.data(), description); + + CHECK(dst == std::vector {6, 15, 24, 33}); +} + +TEST_CASE("reduce bitwise (CPU)") { + // Reduce a (3, 2) row-major input down to 3 outputs, folding over the trailing axis. + std::vector src = { + 0b1100, + 0b1010, // + 0b0110, + 0b0011, // + 0b1111, + 0b0101 // + }; + std::vector dst(3, 0); + + auto description = [&](ReductionOp op) { + ReductionDescription d(DataType::Int32, op); + d.add_dimension(3, 2 * sizeof(int32_t), sizeof(int32_t)); + d.reduction_extent = 2; + d.reduction_stride = sizeof(int32_t); + return d; + }; + + SECTION("BitwiseAnd") { + reduce(src.data(), dst.data(), description(ReductionOp::BitwiseAnd)); + CHECK(dst == std::vector {0b1000, 0b0010, 0b0101}); + } + + SECTION("BitwiseOr") { + reduce(src.data(), dst.data(), description(ReductionOp::BitwiseOr)); + CHECK(dst == std::vector {0b1110, 0b0111, 0b1111}); + } + + SECTION("rejects a floating-point data type") { + auto d = description(ReductionOp::BitwiseAnd); + d.dtype = DataType::Float32; + CHECK_THROWS(reduce(src.data(), dst.data(), d)); + } +} + +TEST_CASE("reduce key-value / argmax (CPU)") { + // Fold 4 (key, value) pairs down to a single pair. + std::vector> src = { + {0, 3.0}, + {1, 7.0}, + {2, 1.0}, + {3, 5.0}, + }; + std::vector> dst(1); + + ReductionDescription description(DataType::KeyValueFloat64, ReductionOp::Max); + description.reduction_extent = 4; + description.reduction_stride = sizeof(KeyValue); + + SECTION("Max keeps the largest value with its key") { + reduce(src.data(), dst.data(), description); + CHECK(dst[0] == KeyValue {1, 7.0}); + } + + SECTION("Min keeps the smallest value with its key") { + description.operation = ReductionOp::Min; + reduce(src.data(), dst.data(), description); + CHECK(dst[0] == KeyValue {2, 1.0}); + } + + SECTION("rejects an unsupported operator") { + description.operation = ReductionOp::Sum; + CHECK_THROWS(reduce(src.data(), dst.data(), description)); + } +} + +TEST_CASE("ReductionDescription::src_range/dst_range") { + constexpr auto elem = static_cast(sizeof(int32_t)); + + SECTION("only the reduced axis") { + ReductionDescription description; + description.dtype = DataType::Int32; + description.input_offset = 4; + description.output_offset = 8; + description.reduction_extent = 3; + description.reduction_stride = elem; + + // src spans the 3 reduced elements; dst is a single element (the reduced axis does not + // move the output pointer). + CHECK(description.src_range() == Range(4, 4 + 3 * elem)); + CHECK(description.dst_range() == Range(8, 8 + elem)); + } + + SECTION("batch axis plus reduced axis") { + ReductionDescription description; + description.dtype = DataType::Int32; + description.add_dimension(4, 3 * elem, elem); + description.reduction_extent = 3; + description.reduction_stride = elem; + + CHECK(description.src_range() == Range(0, 3 * 3 * elem + 2 * elem + elem)); + CHECK(description.dst_range() == Range(0, 3 * elem + elem)); + } + + SECTION("negative reduction stride") { + ReductionDescription description; + description.dtype = DataType::Int32; + description.add_dimension(4, elem, elem); + description.reduction_extent = 3; + description.reduction_stride = -elem; + + // The reduced axis walks backwards, pulling the lower bound below the base offset. + CHECK(description.src_range() == Range(-2 * elem, 3 * elem + elem)); + CHECK(description.dst_range() == Range(0, 3 * elem + elem)); + } +} diff --git a/test/runtime/test_device_event.cpp b/test/runtime/test_device_event.cpp new file mode 100644 index 00000000..66a7e49e --- /dev/null +++ b/test/runtime/test_device_event.cpp @@ -0,0 +1,98 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/device_event.hpp" + +using namespace kmm; + +TEST_CASE("DeviceEvent null") { + DeviceEvent null_event; + CHECK(null_event.is_null()); + CHECK(null_event == DeviceEvent::null()); + + DeviceEvent event {DeviceStreamId(3), 1}; + CHECK_FALSE(event.is_null()); + CHECK(event.stream() == DeviceStreamId(3)); +} + +TEST_CASE("DeviceEvent ordering on same stream") { + auto stream = DeviceStreamId(0); + DeviceEvent a {stream, 1}; + DeviceEvent b {stream, 2}; + + CHECK(a < b); + CHECK(a.precedes(b)); + CHECK_FALSE(b.precedes(a)); +} + +TEST_CASE("DeviceEventSet construction from a single event") { + SECTION("non-null event is kept") { + DeviceEvent event {DeviceStreamId(1), 1}; + DeviceEventSet set {event}; + + CHECK_FALSE(set.is_empty()); + CHECK(set.contains(event)); + } + + SECTION("null event is dropped, not stored") { + DeviceEventSet set {DeviceEvent::null()}; + + CHECK(set.is_empty()); + + for (const auto& e : set) { + CHECK_FALSE(e.is_null()); + } + } +} + +TEST_CASE("DeviceEventSet::insert") { + SECTION("null events are ignored") { + DeviceEventSet set; + set.insert(DeviceEvent::null()); + + CHECK(set.is_empty()); + } + + SECTION("events on the same stream are collapsed to the latest") { + auto stream = DeviceStreamId(2); + DeviceEvent early {stream, 1}; + DeviceEvent late {stream, 5}; + + DeviceEventSet set; + set.insert(early); + set.insert(late); + + CHECK(set.contains(late)); + CHECK(set.find(stream) == late); + + size_t count = 0; + for (const auto& e : set) { + (void)e; + count++; + } + CHECK(count == 1); + } + + SECTION("events on different streams are kept separately") { + DeviceEvent a {DeviceStreamId(0), 1}; + DeviceEvent b {DeviceStreamId(1), 1}; + + DeviceEventSet set; + set.insert(a); + set.insert(b); + + CHECK(set.contains(a)); + CHECK(set.contains(b)); + } +} + +TEST_CASE("DeviceEventSet::insert of another set never introduces null events") { + DeviceEventSet source; + source.insert(DeviceEvent {DeviceStreamId(0), 1}); + + DeviceEventSet dest; + dest.insert(source); + + for (const auto& e : dest) { + CHECK_FALSE(e.is_null()); + } +} diff --git a/test/runtime/test_identifiers.cpp b/test/runtime/test_identifiers.cpp new file mode 100644 index 00000000..2356525a --- /dev/null +++ b/test/runtime/test_identifiers.cpp @@ -0,0 +1,48 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/identifiers.hpp" + +using namespace kmm; + +TEST_CASE("MemoryId::host and MemoryId::device") { + MemoryId host = MemoryId::host(); + CHECK(host.is_host()); + CHECK_FALSE(host.is_device()); + + MemoryId device = MemoryId::device(DeviceId(2)); + CHECK(device.is_device()); + CHECK_FALSE(device.is_host()); + CHECK(device.as_device() == DeviceId(2)); +} + +TEST_CASE("MemoryId from string: host aliases") { + for (const std::string& name : {"host", "cpu", "HOST", "Cpu"}) { + MemoryId id = name; + CHECK(id.is_host()); + } +} + +TEST_CASE("MemoryId from string: device aliases without index default to device 0") { + for (const std::string& name : {"gpu", "cuda", "hip", "device", "GPU", "Cuda"}) { + MemoryId id = name; + CHECK(id.is_device()); + CHECK(id.as_device() == DeviceId(0)); + } +} + +TEST_CASE("MemoryId from string: device aliases with explicit index") { + CHECK(MemoryId("gpu:0").as_device() == DeviceId(0)); + CHECK(MemoryId("cuda:1").as_device() == DeviceId(1)); + CHECK(MemoryId("device:3").as_device() == DeviceId(3)); + CHECK(MemoryId("CUDA:2").as_device() == DeviceId(2)); +} + +TEST_CASE("MemoryId from string: invalid strings throw") { + CHECK_THROWS_AS(MemoryId("host:0"), std::runtime_error); + CHECK_THROWS_AS(MemoryId("gpu:"), std::runtime_error); + CHECK_THROWS_AS(MemoryId("gpu:x"), std::runtime_error); + CHECK_THROWS_AS(MemoryId("gpu:1x"), std::runtime_error); + CHECK_THROWS_AS(MemoryId("gpu:-1"), std::runtime_error); + CHECK_THROWS_AS(MemoryId(""), std::runtime_error); + CHECK_THROWS_AS(MemoryId("nonsense"), std::runtime_error); +} diff --git a/test/runtime/test_memory_manager.cpp b/test/runtime/test_memory_manager.cpp new file mode 100644 index 00000000..5a4515fa --- /dev/null +++ b/test/runtime/test_memory_manager.cpp @@ -0,0 +1,17 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/runtime/memory_manager.hpp" + +using namespace kmm; + +class A: public reference_count {}; + +class B: public A {}; + +TEST_CASE("MemoryManager") { + // auto x = MemoryManager {nullptr, DeviceStreamRegistry {}}; + // auto y = x.create_buffer(BufferLayout::for_type(), "test"); + // + // auto z = make_refcnt(); + // refcnt_ptr a = z; +} \ No newline at end of file diff --git a/test/runtime/test_stream_manager.cpp b/test/runtime/test_stream_manager.cpp new file mode 100644 index 00000000..9eb4fc82 --- /dev/null +++ b/test/runtime/test_stream_manager.cpp @@ -0,0 +1,3 @@ +#include "catch2/catch_all.hpp" + +TEST_CASE("DeviceStreamManager") {} \ No newline at end of file diff --git a/test/utils/test_bounds.cpp b/test/utils/test_bounds.cpp deleted file mode 100644 index 8093527f..00000000 --- a/test/utils/test_bounds.cpp +++ /dev/null @@ -1,292 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/bounds.hpp" - -using namespace kmm; - -TEST_CASE("Bounds") { - SECTION("constructor") { - Bounds<3> a; - Bounds<3> b = {1, 2, 3}; - Bounds<3> c = {Range {4, 10}, Range {10}, 7}; - Bounds<3> d = {Range {-1, 1}, Range {1, 2}, Range {2, 3}}; - - REQUIRE(a[0] == Range {0, 0}); - REQUIRE(a[1] == Range {0, 0}); - REQUIRE(a[2] == Range {0, 0}); - - REQUIRE(b[0] == Range {0, 1}); - REQUIRE(b[1] == Range {0, 2}); - REQUIRE(b[2] == Range {0, 3}); - - REQUIRE(c[0] == Range {4, 10}); - REQUIRE(c[1] == Range {0, 10}); - REQUIRE(c[2] == Range {0, 7}); - - REQUIRE(d[0] == Range {-1, 1}); - REQUIRE(d[1] == Range {1, 2}); - REQUIRE(d[2] == Range {2, 3}); - } - - SECTION("empty") { - Bounds<3> e = Bounds<3>::empty(); - REQUIRE(e[0] == Range {0, 0}); - REQUIRE(e[1] == Range {0, 0}); - REQUIRE(e[2] == Range {0, 0}); - } - - SECTION("one") { - Bounds<3> f = Bounds<3>::one(); - REQUIRE(f[0] == Range {0, 1}); - REQUIRE(f[1] == Range {0, 1}); - REQUIRE(f[2] == Range {0, 1}); - } - - SECTION("from_bounds") { - auto a = Bounds<3>::from_bounds({1, 2, 3}, {4, 5, 6}); - REQUIRE(a[0] == Range {1, 4}); - REQUIRE(a[1] == Range {2, 5}); - REQUIRE(a[2] == Range {3, 6}); - } - - SECTION("from_offset_size") { - auto a = Bounds<3>::from_offset_size({1, 2, 3}, {4, 5, 6}); - REQUIRE(a[0] == Range {1, 1 + 4}); - REQUIRE(a[1] == Range {2, 2 + 5}); - REQUIRE(a[2] == Range {3, 3 + 6}); - } - - SECTION("get_or_default") { - Bounds<3> a = {Range {4, 10}, Range {10}, 7}; - - REQUIRE(a.get_or_default(0) == Range {4, 10}); - REQUIRE(a.get_or_default(1) == Range {0, 10}); - REQUIRE(a.get_or_default(2) == Range {0, 7}); - REQUIRE(a.get_or_default(3) == Range {0, 1}); - } - - SECTION("begin/end/size") { - Bounds<3> a = {Range {4, 10}, Range {10}, 7}; - - REQUIRE(a.begin(0) == 4); - REQUIRE(a.begin(1) == 0); - REQUIRE(a.begin(2) == 0); - REQUIRE(a.begin(3) == 0); - - REQUIRE(a.end(0) == 10); - REQUIRE(a.end(1) == 10); - REQUIRE(a.end(2) == 7); - REQUIRE(a.end(3) == 1); - - REQUIRE(a.size(0) == 6); - REQUIRE(a.size(1) == 10); - REQUIRE(a.size(2) == 7); - REQUIRE(a.size(3) == 1); - - REQUIRE(a.begin() == Point {4, 0, 0}); - REQUIRE(a.end() == Point {10, 10, 7}); - REQUIRE(a.size() == Point {6, 10, 7}); - } - - SECTION("is_empty/volume") { - Bounds<3> a = {1, 1, 1}; - Bounds<3> b = {1, 2, 3}; - Bounds<3> c = {Range {4, 10}, Range {20, 10}, 7}; - Bounds<3> d = {Range {-1, 1}, Range {1, 2}, Range {2, 3}}; - - REQUIRE_FALSE(a.is_empty()); - REQUIRE_FALSE(b.is_empty()); - REQUIRE(c.is_empty()); - REQUIRE_FALSE(d.is_empty()); - - REQUIRE(a.volume() == 1); - REQUIRE(b.volume() == 6); - REQUIRE(c.volume() == 0); - REQUIRE(d.volume() == 2); - } - - SECTION("intersection/unite/overlaps/contains") { - Bounds<3> a = {1, 1, 1}; - Bounds<3> b = {1, 2, 3}; - Bounds<3> c = {Range {4, 10}, Range {20, 10}, 7}; - Bounds<3> d = {Range {-1, 1}, Range {1, 2}, Range {2, 3}}; - -#define REQUIRE_COMMUTATIVE(X, Y) \ - REQUIRE((X).intersection((Y)) == (Y).intersection(X)); \ - REQUIRE((X).unite((Y)) == (Y).unite(X)); \ - REQUIRE((X).overlaps((Y)) == (Y).overlaps(X)); - - REQUIRE_COMMUTATIVE(a, a) - REQUIRE_COMMUTATIVE(a, b) - REQUIRE_COMMUTATIVE(a, c) - REQUIRE_COMMUTATIVE(a, d) - REQUIRE_COMMUTATIVE(b, a) - REQUIRE_COMMUTATIVE(b, b) - REQUIRE_COMMUTATIVE(b, c) - REQUIRE_COMMUTATIVE(b, d) - REQUIRE_COMMUTATIVE(c, a) - REQUIRE_COMMUTATIVE(c, b) - REQUIRE_COMMUTATIVE(c, c) - REQUIRE_COMMUTATIVE(c, d) - REQUIRE_COMMUTATIVE(d, a) - REQUIRE_COMMUTATIVE(d, b) - REQUIRE_COMMUTATIVE(d, c) - REQUIRE_COMMUTATIVE(d, d) - - REQUIRE(a.intersection(a) == a); - REQUIRE(a.intersection(b) == a); - REQUIRE(a.intersection(c) == Bounds {Range {4, 1}, Range {20, 1}, 1}); - REQUIRE(a.intersection(d) == Bounds {1, Range {1, 1}, Range {2, 1}}); - REQUIRE(b.intersection(a) == a); - REQUIRE(b.intersection(b) == b); - REQUIRE(b.intersection(c) == Bounds {Range {4, 1}, Range {20, 2}, 3}); - REQUIRE(b.intersection(d) == Bounds {1, Range {1, 2}, Range {2, 3}}); - REQUIRE(c.intersection(a) == Bounds {Range {4, 1}, Range {20, 1}, 1}); - REQUIRE(c.intersection(b) == Bounds {Range {4, 1}, Range {20, 2}, 3}); - REQUIRE(c.intersection(c) == c); - REQUIRE(c.intersection(d) == Bounds {Range {4, 1}, Range {20, 2}, Range {2, 3}}); - REQUIRE(d.intersection(a) == Bounds {1, Range {1, 1}, Range {2, 1}}); - REQUIRE(d.intersection(b) == Bounds {1, Range {1, 2}, Range {2, 3}}); - REQUIRE(d.intersection(c) == Bounds {Range {4, 1}, Range {20, 2}, Range {2, 3}}); - REQUIRE(d.intersection(d) == d); - - REQUIRE(a.unite(a) == a); - REQUIRE(a.unite(b) == b); - REQUIRE(a.unite(c) == Bounds {10, 10, 7}); - REQUIRE(a.unite(d) == Bounds {Range {-1, 1}, 2, 3}); - REQUIRE(b.unite(a) == b); - REQUIRE(b.unite(b) == b); - REQUIRE(b.unite(c) == Bounds {10, 10, 7}); - REQUIRE(b.unite(d) == Bounds {Range {-1, 1}, 2, 3}); - REQUIRE(c.unite(a) == Bounds {10, 10, 7}); - REQUIRE(c.unite(b) == Bounds {10, 10, 7}); - REQUIRE(c.unite(c) == c); - REQUIRE(c.unite(d) == Bounds {Range {-1, 10}, Range {1, 10}, 7}); - REQUIRE(d.unite(a) == Bounds {Range {-1, 1}, 2, 3}); - REQUIRE(d.unite(b) == Bounds {Range {-1, 1}, 2, 3}); - REQUIRE(d.unite(c) == Bounds {Range {-1, 10}, Range {1, 10}, 7}); - REQUIRE(d.unite(d) == d); - - REQUIRE(a.overlaps(a)); - REQUIRE(a.overlaps(b)); - REQUIRE_FALSE(a.overlaps(c)); - REQUIRE_FALSE(a.overlaps(d)); - REQUIRE(b.overlaps(a)); - REQUIRE(b.overlaps(b)); - REQUIRE_FALSE(b.overlaps(c)); - REQUIRE(b.overlaps(d)); - REQUIRE_FALSE(c.overlaps(a)); - REQUIRE_FALSE(c.overlaps(b)); - REQUIRE_FALSE(c.overlaps(c)); - REQUIRE_FALSE(c.overlaps(d)); - REQUIRE_FALSE(d.overlaps(a)); - REQUIRE(d.overlaps(b)); - REQUIRE_FALSE(d.overlaps(c)); - REQUIRE(d.overlaps(d)); - - REQUIRE(a.contains(a)); - REQUIRE_FALSE(a.contains(b)); - REQUIRE(a.contains(c)); - REQUIRE_FALSE(a.contains(d)); - REQUIRE(b.contains(a)); - REQUIRE(b.contains(b)); - REQUIRE(b.contains(c)); - REQUIRE_FALSE(b.contains(d)); - REQUIRE_FALSE(c.contains(a)); - REQUIRE_FALSE(c.contains(b)); - REQUIRE(c.contains(c)); - REQUIRE_FALSE(c.contains(d)); - REQUIRE_FALSE(d.contains(a)); - REQUIRE_FALSE(d.contains(b)); - REQUIRE(d.contains(c)); - REQUIRE(d.contains(d)); - } - - SECTION("contains") { - auto simple = Bounds {10, Range {10, 20}, 30}; - REQUIRE_FALSE(simple.contains(0, 0, 0)); - REQUIRE(simple.contains(0, 10, 10)); - REQUIRE(simple.contains(0, 15, 0)); - REQUIRE(simple.contains(9, 19, 29)); - REQUIRE_FALSE(simple.contains(10, 20, 30)); - REQUIRE_FALSE(simple.contains(-1, 0, 0)); - - auto empty = Bounds {10, Range {20, 20}, 30}; - REQUIRE_FALSE(empty.contains(0, 0, 0)); - REQUIRE_FALSE(empty.contains(0, 10, 10)); - REQUIRE_FALSE(empty.contains(10, 20, 30)); - REQUIRE_FALSE(empty.contains(-1, 0, 0)); - } - - SECTION("shift_by") { - auto x = Bounds {1, 2, 3}; - - auto y = x.shift_by(Point {3, 4, 5}); - REQUIRE(y[0] == Range {3, 3 + 1}); - REQUIRE(y[1] == Range {4, 4 + 2}); - REQUIRE(y[2] == Range {5, 5 + 3}); - - auto z = y.shift_by(Point {3, 4, 5}); - REQUIRE(z[0] == Range {6, 6 + 1}); - REQUIRE(z[1] == Range {8, 8 + 2}); - REQUIRE(z[2] == Range {10, 10 + 3}); - } - - SECTION("split_tail_along") { - auto x = Bounds {10, 20, 30}; - - SECTION("valid") { - auto y = x.split_tail_along(1, 10); - REQUIRE(x == Bounds {10, Range {0, 10}, 30}); - REQUIRE(y == Bounds {10, Range {10, 20}, 30}); - } - - SECTION("invalid axis") { - auto y = x.split_tail_along(10, 1); - REQUIRE(x == Bounds {10, 20, 30}); - REQUIRE(y.is_empty()); - } - - SECTION("below start") { - auto y = x.split_tail_along(1, -10); - REQUIRE(x == Bounds {10, Range {0, 0}, 30}); - REQUIRE(y == Bounds {10, Range {0, 20}, 30}); - } - - SECTION("after end") { - auto y = x.split_tail_along(1, 100); - REQUIRE(x == Bounds {10, Range {0, 20}, 30}); - REQUIRE(y == Bounds {10, Range {20, 20}, 30}); - } - } - - SECTION("concat") { - auto x = Bounds {1, 2, 3}; - auto y = Bounds {4, 5}; - REQUIRE(concat(x, y) == Bounds {1, 2, 3, 4, 5}); - } - - SECTION("operator==") { - auto a = Bounds {3, 2, 1}; - auto b = Bounds {3, 2}; - auto c = Bounds {4, 5, 6}; - - REQUIRE(a == a); - REQUIRE(a == b); - REQUIRE(a != c); - - REQUIRE(b == a); - REQUIRE(b == b); - REQUIRE(b != c); - - REQUIRE(c != a); - REQUIRE(c != b); - REQUIRE(c == c); - } - - SECTION("operator<<") { - std::stringstream stream; - stream << Bounds {1, Range {2, 3}, Range {1, -1}}; - REQUIRE(stream.str() == "{0...1, 2...3, 1...-1}"); - } -} \ No newline at end of file diff --git a/test/utils/test_checked_compare.cpp b/test/utils/test_checked_compare.cpp deleted file mode 100644 index fd8d989c..00000000 --- a/test/utils/test_checked_compare.cpp +++ /dev/null @@ -1,272 +0,0 @@ -#include -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/utils/checked_compare.hpp" - -using namespace kmm; - -using u8 = uint8_t; -using u16 = uint16_t; -using u32 = uint32_t; -using u64 = uint64_t; - -using i8 = int8_t; -using i16 = int16_t; -using i32 = int32_t; -using i64 = int64_t; - -using f32 = float; -using f64 = double; - -#define REQUIRE_IS_LESS(A, B) \ - REQUIRE_FALSE(is_less(A, A)); \ - REQUIRE_FALSE(is_less(B, B)); \ - REQUIRE(is_less(A, B)); \ - REQUIRE_FALSE(is_less(B, A)); - -TEST_CASE("is_less") { - REQUIRE_IS_LESS(i8(-128), u8(0)); - REQUIRE_IS_LESS(u8(0), i8(127)); - REQUIRE_IS_LESS(i16(-1), u16(1)); - REQUIRE_IS_LESS(u16(1), i16(2)); - REQUIRE_IS_LESS(i32(-100), u32(0)); - REQUIRE_IS_LESS(u32(0), i64(1)); - REQUIRE_IS_LESS(i64(std::numeric_limits::min()), u64(0ULL)); - REQUIRE_IS_LESS(u64(0ULL), i64(std::numeric_limits::max())); - REQUIRE_IS_LESS(f32(1.5f), f64(1.6)); - REQUIRE_IS_LESS(f64(-1.0), f32(0.0f)); - REQUIRE_IS_LESS(f32(-std::numeric_limits::max()), f64(std::numeric_limits::max())); - REQUIRE_IS_LESS(i32(std::numeric_limits::min()), i32(0)); - REQUIRE_IS_LESS(i32(-2147483647), u32(2147483648u)); - REQUIRE_IS_LESS(u16(0), f64(0.1)); - REQUIRE_IS_LESS(f64(-std::numeric_limits::infinity()), f64(-1e308)); - REQUIRE_IS_LESS(f32(-1e10f), f64(-1e9)); - REQUIRE_IS_LESS(u64(100), f64(101.0)); - REQUIRE_IS_LESS(i64(-1000), f32(0.0f)); - REQUIRE_IS_LESS(f32(0.999f), i32(1)); - REQUIRE_IS_LESS(i8(-10), f32(-9.9f)); - REQUIRE_IS_LESS(u8(200), u16(300)); - REQUIRE_IS_LESS(i16(30000), i32(40000)); - REQUIRE_IS_LESS(i32(1), f64(1.1)); - REQUIRE_IS_LESS(f64(1.1), f64(1.2L)); - REQUIRE_IS_LESS(f64(-1.2L), f64(-1.1)); - REQUIRE_IS_LESS(f64(0.0L), f64(0.1L)); - REQUIRE_IS_LESS(i64(-1), f64(0.0)); - REQUIRE_IS_LESS(u32(1), f64(1.1)); - REQUIRE_IS_LESS(f64(-0.01), u32(0)); - REQUIRE_IS_LESS(f32(-0.1f), u16(0)); - REQUIRE_IS_LESS(u64(1), f64(1.0001L)); - REQUIRE_IS_LESS(i8(10), f64(10.1L)); - REQUIRE_IS_LESS(i16(123), f64(123.1L)); - REQUIRE_IS_LESS(u16(65535), f64(65535.5)); - REQUIRE_IS_LESS(i32(2147483646), i64(2147483647)); - REQUIRE_IS_LESS(u32(4294967294U), u64(4294967295ULL)); - REQUIRE_IS_LESS(i64(9223372036854775806LL), u64(9223372036854775807ULL)); - REQUIRE_IS_LESS(f64(0.0), std::numeric_limits::infinity()); - REQUIRE_IS_LESS(f32(-std::numeric_limits::infinity()), f32(0.0f)); - REQUIRE_IS_LESS(i32(-1), f64(0.0L)); - REQUIRE_IS_LESS(f32(std::numeric_limits::max()), f64(std::numeric_limits::max())); - REQUIRE_IS_LESS(i32(-1), u32(10)); - REQUIRE_IS_LESS(f64(0.0), std::numeric_limits::infinity()); - REQUIRE_IS_LESS(i64(0), u64(std::numeric_limits::max())); - REQUIRE_IS_LESS(f32(1.0001f), f64(1.0002)); - REQUIRE_IS_LESS(i32(-1), i32(0)); - - REQUIRE_FALSE(is_less(-0.0, 0.0)); - REQUIRE_FALSE(is_less(f32(1.1f), f64(1.1))); - REQUIRE_FALSE(is_less(u8(255), u8(255))); - REQUIRE_FALSE(is_less(std::numeric_limits::quiet_NaN(), f64(0.0))); -} - -#define REQUIRE_IS_EQUAL(A, B) \ - REQUIRE(is_equal(A, B)); \ - REQUIRE_FALSE(is_less(A, B)); \ - REQUIRE_FALSE(is_less(B, A)); - -#define REQUIRE_NOT_EQUAL(A, B) \ - REQUIRE_FALSE(is_equal(A, B)); \ - REQUIRE((is_less(A, B) || is_less(B, A))); - -TEST_CASE("is_equal") { - REQUIRE_IS_EQUAL(i8(0), u8(0)); - REQUIRE_IS_EQUAL(i16(123), u16(123)); - REQUIRE_IS_EQUAL(i32(-100), i64(-100)); - REQUIRE_IS_EQUAL(u32(3000), i64(3000)); - REQUIRE_IS_EQUAL(f32(2.5f), f64(2.5)); - REQUIRE_IS_EQUAL(f64(-1.0), f64(-1.0L)); - REQUIRE_IS_EQUAL(u64(0ULL), i64(0LL)); - REQUIRE_IS_EQUAL(i32(std::numeric_limits::max()), i64(std::numeric_limits::max())); - REQUIRE_IS_EQUAL( - f32(std::numeric_limits::infinity()), - f64(std::numeric_limits::infinity()) - ); - REQUIRE_IS_EQUAL(f64(-0.0), f32(-0.0f)); - REQUIRE_IS_EQUAL(i8(-128), i16(-128)); - REQUIRE_IS_EQUAL(u16(65535), u32(65535)); - REQUIRE_IS_EQUAL(i64(-1), f64(-1.0L)); - REQUIRE_IS_EQUAL(f64(0.0L), f64(0.0)); - REQUIRE_IS_EQUAL(u8(255), i32(255)); - REQUIRE_IS_EQUAL(i32(2147483647), f64(2147483647.0)); - REQUIRE_IS_EQUAL(u32(1), f64(1.0)); - REQUIRE_IS_EQUAL(f32(0.5f), f64(0.5L)); - REQUIRE_IS_EQUAL((unsigned long long)(42), f64(42.0L)); - REQUIRE_IS_EQUAL(i16(-32768), i32(-32768)); - REQUIRE_IS_EQUAL(i64(1234567890123LL), u64(1234567890123ULL)); - REQUIRE_IS_EQUAL(u32(0), f32(0.0f)); - REQUIRE_IS_EQUAL( - f64(-std::numeric_limits::infinity()), - f64(-std::numeric_limits::infinity()) - ); - REQUIRE_IS_EQUAL(f64(123456.0L), f64(123456.0L)); - - REQUIRE_NOT_EQUAL(i8(0), u8(1)); - REQUIRE_NOT_EQUAL(i16(123), u16(124)); - REQUIRE_NOT_EQUAL(i32(-100), i64(100)); - REQUIRE_NOT_EQUAL(u32(3000), i64(-3000)); - REQUIRE_NOT_EQUAL(f32(2.5f), f64(2.5001)); - REQUIRE_NOT_EQUAL(f64(-1.0), f64(-1.0000000001L)); - REQUIRE_NOT_EQUAL(u64(1), i64(2)); - REQUIRE_NOT_EQUAL( - i32(std::numeric_limits::max()), - i64(std::numeric_limits::max()) - 1 - ); - REQUIRE_NOT_EQUAL( - u64(std::numeric_limits::max()), - f64(std::numeric_limits::max()) - 1.0 - ); - REQUIRE_NOT_EQUAL(u64(std::numeric_limits::max()), f64(std::numeric_limits::max())); - REQUIRE_NOT_EQUAL( - f32(std::numeric_limits::infinity()), - f64(std::numeric_limits::lowest()) - ); - REQUIRE_NOT_EQUAL(f64(123.0), f64(124.0)); - REQUIRE_NOT_EQUAL(f64(1.0L), f64(1.0000000001L)); - REQUIRE_NOT_EQUAL(i8(-128), i16(-127)); - REQUIRE_NOT_EQUAL(u16(65534), u32(65535)); - REQUIRE_NOT_EQUAL(i64(-1), f64(1.0L)); - REQUIRE_NOT_EQUAL(f64(0.0L), f64(0.1)); - REQUIRE_NOT_EQUAL(u8(255), i32(254)); - REQUIRE_NOT_EQUAL(i32(2147483647), f64(2147483646.0)); - REQUIRE_NOT_EQUAL(u32(1), f64(2.0)); - REQUIRE_NOT_EQUAL(f32(0.5f), f64(0.5000001L)); - REQUIRE_NOT_EQUAL((unsigned long long)(42), f64(42.1)); - REQUIRE_NOT_EQUAL(i16(-32768), i32(-32767)); - REQUIRE_NOT_EQUAL(i64(1234567890123LL), u64(1234567890124ULL)); - REQUIRE_NOT_EQUAL(u32(0), f32(0.0001f)); - - REQUIRE_FALSE(is_equal(std::numeric_limits::quiet_NaN(), f64(0.0))); - REQUIRE_FALSE(is_equal(std::numeric_limits::quiet_NaN(), i32(0))); - REQUIRE_FALSE( - is_equal(std::numeric_limits::quiet_NaN(), std::numeric_limits::quiet_NaN()) - ); -} - -TEST_CASE("is_convertible") { - REQUIRE(is_convertible(f64(0.0))); - REQUIRE(is_convertible(f64(2147483647.0))); - REQUIRE(is_convertible(i32(0))); - REQUIRE(is_convertible(i64(4294967295LL))); - REQUIRE(is_convertible(i8(127))); - REQUIRE(is_convertible(i32(255))); - REQUIRE(is_convertible(f64(9007199254740992.0))); - REQUIRE(is_convertible(i32(123456))); - REQUIRE(is_convertible(f32(1.5f))); - REQUIRE(is_convertible(f64(1.5))); - REQUIRE(is_convertible(f64(0.0))); - REQUIRE(is_convertible(f32(-100.0f))); - REQUIRE(is_convertible(i16(-128))); - REQUIRE(is_convertible(u32(65535))); - REQUIRE(is_convertible(i32(std::numeric_limits::min()))); - REQUIRE(is_convertible(f64(1.0L))); - REQUIRE(is_convertible(std::numeric_limits::min())); - REQUIRE(is_convertible(f64(-2147483648.0))); - REQUIRE(is_convertible(u32(100))); - REQUIRE(is_convertible(f64(9007199254740991.0))); - REQUIRE(is_convertible(f32(16777216.0f))); - REQUIRE(is_convertible(f64(0.0))); - REQUIRE(is_convertible(f64(-32768.0))); - REQUIRE(is_convertible(u64(12345ULL))); - REQUIRE(is_convertible(i64(0))); - REQUIRE(is_convertible(std::numeric_limits::infinity())); - REQUIRE(is_convertible(1.0 + std::numeric_limits::epsilon())); - - REQUIRE_FALSE(is_convertible(f64(0.1))); - REQUIRE_FALSE(is_convertible(f64(2147483648.0))); - REQUIRE_FALSE(is_convertible(i32(-1))); - REQUIRE_FALSE(is_convertible(i16(128))); - REQUIRE_FALSE(is_convertible(f64(255.5))); - REQUIRE_FALSE(is_convertible(f64(40000.0))); - REQUIRE_FALSE(is_convertible(f64(std::numeric_limits::infinity()))); - REQUIRE_FALSE(is_convertible(f64(std::numeric_limits::quiet_NaN()))); - REQUIRE_FALSE(is_convertible(f64(9.223372036854775808e18))); - REQUIRE_FALSE(is_convertible(f64(-0.1))); - REQUIRE_FALSE(is_convertible(i32(70000))); - REQUIRE_FALSE(is_convertible(f32(2147483648.0f))); - REQUIRE_FALSE(is_convertible(f64(-32769.0))); - REQUIRE_FALSE(is_convertible(i32(256))); - REQUIRE_FALSE(is_convertible(i32(-1))); - REQUIRE_FALSE(is_convertible(f64(-1.0))); - REQUIRE_FALSE(is_convertible(f64(1.5L))); - REQUIRE_FALSE(is_convertible(f64(std::numeric_limits::max()))); - REQUIRE_FALSE(is_convertible(f64(-2147483649.0))); - REQUIRE_FALSE(is_convertible(f64(std::numeric_limits::infinity()))); - REQUIRE_FALSE(is_convertible(f32(1.0000001f))); - REQUIRE_FALSE(is_convertible(f64(-0.00001))); - REQUIRE_FALSE(is_convertible(f64(std::numeric_limits::quiet_NaN()))); - REQUIRE_FALSE(is_convertible(f64(1e20))); - REQUIRE_FALSE(is_convertible(f32(32768.0f))); - REQUIRE_FALSE(is_convertible(1.0 + std::numeric_limits::epsilon())); - REQUIRE_FALSE(is_convertible(1.0f + std::numeric_limits::epsilon())); - REQUIRE_FALSE(is_convertible(std::numeric_limits::epsilon())); -} - -TEST_CASE("in_range") { - REQUIRE(in_range(u32(5), i32(10))); - REQUIRE(in_range(u32(8), std::numeric_limits::infinity())); - REQUIRE(in_range(u32(5), f64(100.0))); - REQUIRE(in_range(u32(0), i32(1))); - REQUIRE(in_range(i64(72), u8(255))); - REQUIRE(in_range(u32(100), f64(100.000001))); - REQUIRE(in_range(std::numeric_limits::max(), std::numeric_limits::max())); - - REQUIRE_FALSE(in_range(u32(5), i32(-10))); - REQUIRE_FALSE(in_range(std::numeric_limits::infinity(), i8(72))); - REQUIRE_FALSE(in_range(std::numeric_limits::quiet_NaN(), 100)); - REQUIRE_FALSE(in_range(u32(100), f64(99.99))); - REQUIRE_FALSE(in_range(u32(100), f64(99.99))); - REQUIRE_FALSE(in_range(u32(1), i32(0))); - REQUIRE_FALSE(in_range(u8(255), i64(72))); - REQUIRE_FALSE(in_range(std::numeric_limits::max(), std::numeric_limits::max())); -} - -TEST_CASE("checked_cast") { - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - - REQUIRE_NOTHROW(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - REQUIRE_THROWS(checked_cast(std::numeric_limits::max())); - - REQUIRE_THROWS(checked_cast(i32(200))); - REQUIRE_THROWS(checked_cast(f64(200.0))); - REQUIRE_THROWS(checked_cast(i64(-1337))); - - REQUIRE_NOTHROW(checked_cast(i32(200))); - REQUIRE_NOTHROW(checked_cast(f64(201.0))); - REQUIRE_NOTHROW(checked_cast(i32(202))); -} \ No newline at end of file diff --git a/test/utils/test_checked_math.cpp b/test/utils/test_checked_math.cpp deleted file mode 100644 index 5b95738b..00000000 --- a/test/utils/test_checked_math.cpp +++ /dev/null @@ -1,101 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/checked_math.hpp" - -using namespace kmm; - -using i8 = int8_t; -using i32 = int32_t; -using u32 = uint32_t; - -TEST_CASE("checked_add") { - REQUIRE(checked_add(i32(5), i32(1)) == i32(6)); - REQUIRE(checked_add(i32(-5), i32(1)) == i32(-4)); - REQUIRE_THROWS(checked_add(std::numeric_limits::max(), i32(1))); - REQUIRE_THROWS(checked_add(std::numeric_limits::min(), i32(-1))); - - REQUIRE(checked_add(u32(5), u32(1)) == u32(6)); - REQUIRE(checked_add(u32(0), u32(1)) == u32(1)); - REQUIRE_THROWS(checked_add(std::numeric_limits::max(), u32(1))); - - REQUIRE(checked_add(i8(5), i8(1)) == i8(6)); - REQUIRE(checked_add(i8(-5), i8(1)) == i8(-4)); - REQUIRE_THROWS(checked_add(std::numeric_limits::max(), i8(1))); - REQUIRE_THROWS(checked_add(std::numeric_limits::min(), i8(-1))); -} - -TEST_CASE("checked_sub") { - REQUIRE(checked_sub(i32(5), i32(1)) == i32(4)); - REQUIRE(checked_sub(i32(-5), i32(1)) == i32(-6)); - REQUIRE_THROWS(checked_sub(std::numeric_limits::max(), i32(-1))); - REQUIRE_THROWS(checked_sub(std::numeric_limits::min(), i32(1))); - - REQUIRE(checked_sub(u32(5), u32(1)) == u32(4)); - REQUIRE_THROWS(checked_sub(u32(0), u32(1))); - REQUIRE(checked_sub(std::numeric_limits::max(), std::numeric_limits::max()) == 0); - - REQUIRE(checked_sub(i8(5), i8(1)) == i8(4)); - REQUIRE(checked_sub(i8(-5), i8(1)) == i8(-6)); - REQUIRE_THROWS(checked_sub(std::numeric_limits::max(), i8(-1))); - REQUIRE_THROWS(checked_sub(std::numeric_limits::min(), i8(1))); -} - -TEST_CASE("checked_mul") { - REQUIRE(checked_mul(i32(5), i32(1)) == i32(5)); - REQUIRE(checked_mul(i32(-5), i32(1)) == i32(-5)); - REQUIRE(checked_mul(std::numeric_limits::max(), i32(-1)) == -2147483647); - REQUIRE_THROWS(checked_mul(std::numeric_limits::min(), i32(-1))); - - REQUIRE(checked_mul(u32(5), u32(1)) == u32(5)); - REQUIRE(checked_mul(u32(0), u32(1)) == u32(0)); - REQUIRE_THROWS(checked_mul(std::numeric_limits::max(), std::numeric_limits::max())); - - REQUIRE(checked_mul(i8(5), i8(-1)) == i8(-5)); - REQUIRE(checked_mul(i8(-5), i8(-1)) == i8(5)); - REQUIRE(checked_mul(std::numeric_limits::max(), i8(-1)) == -127); - REQUIRE_THROWS(checked_mul(std::numeric_limits::min(), i8(-1))); -} - -TEST_CASE("checked_sum") { - SECTION("empty list") { - std::array values = {}; - REQUIRE(checked_sum(values.data(), values.data()) == 0); - } - - SECTION("no overflow") { - std::array values = {1, 2, 3, 4}; - REQUIRE(checked_sum(values.data(), values.data() + 4) == 10); - } - - SECTION("overflows") { - std::array values = {1, 2, std::numeric_limits::max(), 4}; - REQUIRE_THROWS(checked_sum(values.data(), values.data() + 4)); - } - - SECTION("no overflow with upcast") { - std::array values = {1, 2, std::numeric_limits::max(), 4}; - REQUIRE(checked_sum(values.data(), values.data() + 4, uint64_t(1)) == 2147483655LL); - } -} - -TEST_CASE("checked_product") { - SECTION("empty list") { - std::array values = {}; - REQUIRE(checked_product(values.data(), values.data()) == 1); - } - - SECTION("no overflow") { - std::array values = {1, 2, 3, 4}; - REQUIRE(checked_product(values.data(), values.data() + 4) == 24); - } - - SECTION("overflows") { - std::array values = {1, std::numeric_limits::max() / 2, 3}; - REQUIRE_THROWS(checked_product(&*values.begin(), &*values.end())); - } - - SECTION("no overflow with upcast") { - std::array values = {1, std::numeric_limits::max() / 2, 3}; - REQUIRE(checked_product(&*values.begin(), &*values.end(), uint64_t(1)) == 3221225469LL); - } -} \ No newline at end of file diff --git a/test/utils/test_dim.cpp b/test/utils/test_dim.cpp deleted file mode 100644 index 64c0a1e5..00000000 --- a/test/utils/test_dim.cpp +++ /dev/null @@ -1,223 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/dim.hpp" - -using namespace kmm; - -TEST_CASE("Dim") { - SECTION("constructor") { - Dim<3> a; - Dim<3> b = {-1, 1, 2}; - - REQUIRE(a[0] == 1); - REQUIRE(a[1] == 1); - REQUIRE(a[2] == 1); - - REQUIRE(b[0] == -1); - REQUIRE(b[1] == 1); - REQUIRE(b[2] == 2); - } - - SECTION("conversions") { - Dim<3> a = {1, 2, 3}; - Dim<3> b = {-1, 1, 1}; - - REQUIRE_FALSE(a.is_convertible_to<2>()); - REQUIRE(a.is_convertible_to<4>()); - REQUIRE_FALSE(a.is_convertible_to<2, unsigned int>()); - REQUIRE(a.is_convertible_to<3, unsigned int>()); - REQUIRE(a.is_convertible_to<4, unsigned int>()); - - REQUIRE(b.is_convertible_to<2>()); - REQUIRE(b.is_convertible_to<4>()); - REQUIRE_FALSE(b.is_convertible_to<2, unsigned int>()); - REQUIRE_FALSE(b.is_convertible_to<3, unsigned int>()); - REQUIRE_FALSE(b.is_convertible_to<4, unsigned int>()); - - REQUIRE_THROWS(Dim<2> {a}); - REQUIRE(Dim<4> {a} == Dim<4> {1, 2, 3, 1}); - REQUIRE_THROWS(Dim<2, unsigned int> {a}); - REQUIRE(Dim<4, unsigned int> {a} == Dim<4, unsigned int> {1, 2, 3, 1}); - - REQUIRE(Dim<2> {b} == Dim<2> {-1, 1}); - REQUIRE(Dim<4> {b} == Dim<4> {-1, 1, 1, 1}); - REQUIRE_THROWS(Dim<2, unsigned int> {b}); - REQUIRE_THROWS(Dim<4, unsigned int> {b}); - - REQUIRE(Dim<2>::from(a) == Dim<2> {1, 2}); - REQUIRE(Dim<4>::from(a) == Dim<4> {1, 2, 3, 1}); - REQUIRE(Dim<2, unsigned int>::from(a) == Dim<2, unsigned int> {1, 2}); - REQUIRE(Dim<4, unsigned int>::from(a) == Dim<4, unsigned int> {1, 2, 3, 1}); - - REQUIRE(Dim<2>::from(b) == Dim<2> {-1, 1}); - REQUIRE(Dim<4>::from(b) == Dim<4> {-1, 1, 1, 1}); - REQUIRE(Dim<2, unsigned int>::from(b) == Dim<2, unsigned int> {UINT_MAX, 1}); - REQUIRE(Dim<4, unsigned int>::from(b) == Dim<4, unsigned int> {UINT_MAX, 1, 1, 1}); - } - - SECTION("fill") { - auto fill = Dim<3>::fill(1337); - auto one = Dim<3>::one(); - auto zero = Dim<3>::zero(); - - REQUIRE(fill == Dim<3> {1337, 1337, 1337}); - REQUIRE(one == Dim<3> {1, 1, 1}); - REQUIRE(zero == Dim<3> {0, 0, 0}); - } - - SECTION("get_or_default") { - Dim<3> a; - Dim<3> b = {-1, 1, 2}; - - REQUIRE(a.get_or_default(0) == 1); - REQUIRE(a.get_or_default(1) == 1); - REQUIRE(a.get_or_default(2) == 1); - REQUIRE(a.get_or_default(3) == 1); - REQUIRE(a.get_or_default(3, 1337) == 1337); - - REQUIRE(b.get_or_default(0) == -1); - REQUIRE(b.get_or_default(1) == 1); - REQUIRE(b.get_or_default(2) == 2); - REQUIRE(b.get_or_default(3) == 1); - REQUIRE(b.get_or_default(3, 1337) == 1337); - } - - SECTION("is_empty") { - Dim<3> a = {1, 2, 3}; - Dim<3> b = {1, -1, 1}; - Dim<3> c = {1, 0, 1}; - - REQUIRE(a.is_empty() == false); - REQUIRE(b.is_empty()); - REQUIRE(c.is_empty()); - } - - SECTION("is_empty") { - Dim<3> a = {1, 2, 3}; - Dim<3> b = {1, -1, 1}; - Dim<3> c = {1, 0, 1}; - - REQUIRE(a.volume() == 6); - REQUIRE(b.volume() == 0); - REQUIRE(c.volume() == 0); - } - - SECTION("contains") { - Dim<3> a = {0, 0, 0}; - Dim<3> b = {1, 1, 1}; - Dim<3> c = {INT64_MAX, INT64_MAX, INT64_MAX}; - Dim<3> d = {1, -1, 1}; - - Point<3> p0 = {0, 0, 0}; - Point<3> p1 = {1, 0, 0}; - Point<3> p2 = {INT64_MAX, INT64_MAX, INT64_MAX}; - Point<3> p3 = {INT64_MAX - 1, INT64_MAX - 1, INT64_MAX - 1}; - Point<2> p4 = {1, 0}; - Point<5> p5 = {1, 2, 3, 0, 0}; - Point<3, unsigned long> p6 = {1, 2, 3}; - - REQUIRE_FALSE(a.contains(p0)); - REQUIRE_FALSE(a.contains(p1)); - REQUIRE_FALSE(a.contains(p2)); - REQUIRE_FALSE(a.contains(p3)); - REQUIRE_FALSE(a.contains(p4)); - REQUIRE_FALSE(a.contains(p5)); - REQUIRE_FALSE(a.contains(p6)); - - REQUIRE(b.contains(p0)); - REQUIRE_FALSE(b.contains(p1)); - REQUIRE_FALSE(b.contains(p2)); - REQUIRE_FALSE(b.contains(p3)); - REQUIRE_FALSE(b.contains(p4)); - REQUIRE_FALSE(b.contains(p5)); - REQUIRE_FALSE(b.contains(p6)); - - REQUIRE(c.contains(p0)); - REQUIRE(c.contains(p1)); - REQUIRE_FALSE(c.contains(p2)); - REQUIRE(c.contains(p3)); - REQUIRE(c.contains(p4)); - REQUIRE(c.contains(p5)); - REQUIRE(c.contains(p6)); - - REQUIRE_FALSE(d.contains(p0)); - REQUIRE_FALSE(d.contains(p1)); - REQUIRE_FALSE(d.contains(p2)); - REQUIRE_FALSE(d.contains(p3)); - REQUIRE_FALSE(d.contains(p4)); - REQUIRE_FALSE(d.contains(p5)); - REQUIRE_FALSE(d.contains(p6)); - } - - SECTION("concat") { - Dim<2> a = {1, 2}; - Dim<3> b = {3, 4, 5}; - Dim<5> c = concat(a, b); - - REQUIRE(c == Dim<5> {1, 2, 3, 4, 5}); - } - - SECTION("operator==") { - Dim<3> a = {-1, 2, 3}; - Dim<2> b = {-1, 2}; - Dim<3> c = {-1, 2, 1}; - Dim<3, uint64_t> d = {uint64_t(-1), 2, 3}; - Dim<3, double> e = {-1, 2, 3}; - - REQUIRE(a == a); - REQUIRE(a != b); - REQUIRE(a != c); - REQUIRE(a != d); - REQUIRE(a == e); - - REQUIRE(b != a); - REQUIRE(b == b); - REQUIRE(b == c); - REQUIRE(b != d); - REQUIRE(b != e); - - REQUIRE(c != a); - REQUIRE(c == b); - REQUIRE(c == c); - REQUIRE(c != d); - REQUIRE(c != e); - - REQUIRE(d != a); - REQUIRE(d != b); - REQUIRE(d != c); - REQUIRE(d == d); - REQUIRE(d != e); - - REQUIRE(e == a); - REQUIRE(e != b); - REQUIRE(e != c); - REQUIRE(e != d); - REQUIRE(e == e); - } - - SECTION("is_less") { - REQUIRE(is_less(Dim(int(1)), int(2))); - REQUIRE(is_less(int(1), Dim(int(2)))); - REQUIRE(is_less(Dim(int(1)), Dim(int(2)))); - - REQUIRE(is_less(Dim(int(1)), uint(2))); - REQUIRE(is_less(int(1), Dim(uint(2)))); - REQUIRE(is_less(Dim(int(1)), Dim(uint(2)))); - } - - SECTION("checked_cast") { - REQUIRE(checked_cast(Dim(int(1))) == 1); - REQUIRE(checked_cast>(Dim(int(1))) == Dim(1)); - REQUIRE(checked_cast>(int(1)) == Dim(1)); - - REQUIRE(checked_cast(Dim(uint(1))) == 1); - REQUIRE(checked_cast>(Dim(uint(1))) == Dim(1)); - REQUIRE(checked_cast>(uint(1)) == Dim(1)); - } - - SECTION("operator<<") { - std::stringstream stream; - stream << Dim {1, 2, 3}; - REQUIRE(stream.str() == "{1, 2, 3}"); - } -} \ No newline at end of file diff --git a/test/utils/test_fixed_vector.cpp b/test/utils/test_fixed_vector.cpp deleted file mode 100644 index a09ebdac..00000000 --- a/test/utils/test_fixed_vector.cpp +++ /dev/null @@ -1,159 +0,0 @@ -#include -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/utils/fixed_vector.hpp" - -using namespace kmm; - -TEST_CASE("fixed_vector") { - SECTION("size/alignment") { - REQUIRE(sizeof(fixed_vector) == 1); - REQUIRE(sizeof(fixed_vector) == 4); - REQUIRE(sizeof(fixed_vector) == 8); - REQUIRE(sizeof(fixed_vector) == 16); - REQUIRE(sizeof(fixed_vector) == 16); - REQUIRE(sizeof(fixed_vector) == 32); - REQUIRE(sizeof(fixed_vector) == 32); - REQUIRE(sizeof(fixed_vector) == 32); - REQUIRE(sizeof(fixed_vector) == 32); - - REQUIRE(alignof(fixed_vector) == 1); - REQUIRE(alignof(fixed_vector) == 4); - REQUIRE(alignof(fixed_vector) == 8); - REQUIRE(alignof(fixed_vector) == 16); - REQUIRE(alignof(fixed_vector) == 16); - REQUIRE(alignof(fixed_vector) == 16); - REQUIRE(alignof(fixed_vector) == 16); - REQUIRE(alignof(fixed_vector) == 16); - REQUIRE(alignof(fixed_vector) == 16); - } - - SECTION("N=0") { - fixed_vector a; - fixed_vector b; - - (void)a; - (void)b; - } - - SECTION("N=1") { - fixed_vector a; - REQUIRE(a.x == 0); - - fixed_vector b = {"a"}; - REQUIRE(b.x == "a"); - } - - SECTION("N=2") { - fixed_vector a; - REQUIRE(a.x == 0); - REQUIRE(a.y == 0); - - fixed_vector b = {"a", "b"}; - REQUIRE(b.x == "a"); - REQUIRE(b.y == "b"); - - REQUIRE(b[0] == "a"); - REQUIRE(b[1] == "b"); - } - - SECTION("N=3") { - fixed_vector a; - REQUIRE(a.x == 0); - REQUIRE(a.y == 0); - REQUIRE(a.z == 0); - - fixed_vector b = {"a", "b", "c"}; - REQUIRE(b.x == "a"); - REQUIRE(b.y == "b"); - REQUIRE(b.z == "c"); - - REQUIRE(b[0] == "a"); - REQUIRE(b[1] == "b"); - REQUIRE(b[2] == "c"); - } - - SECTION("N=4") { - fixed_vector a; - REQUIRE(a.x == 0); - REQUIRE(a.y == 0); - REQUIRE(a.z == 0); - REQUIRE(a.w == 0); - - fixed_vector b = {"a", "b", "c", "d"}; - REQUIRE(b.x == "a"); - REQUIRE(b.y == "b"); - REQUIRE(b.z == "c"); - REQUIRE(b.w == "d"); - - REQUIRE(b[0] == "a"); - REQUIRE(b[1] == "b"); - REQUIRE(b[2] == "c"); - REQUIRE(b[3] == "d"); - } - - SECTION("N=5") { - fixed_vector a; - REQUIRE(a[0] == 0); - REQUIRE(a[1] == 0); - REQUIRE(a[2] == 0); - REQUIRE(a[3] == 0); - REQUIRE(a[4] == 0); - - fixed_vector b = {"a", "b", "c", "d", "e"}; - REQUIRE(b[0] == "a"); - REQUIRE(b[1] == "b"); - REQUIRE(b[2] == "c"); - REQUIRE(b[3] == "d"); - REQUIRE(b[4] == "e"); - } - - SECTION("operator==") { - fixed_vector a = {-1, 1}; - fixed_vector b = {-1, 1, 2}; - fixed_vector c = {-1, 1, 2}; - fixed_vector d = {uint(-1), 1, 2}; - - REQUIRE(a == a); - REQUIRE(a != b); - REQUIRE(a != c); - REQUIRE(a != d); - - REQUIRE(b != a); - REQUIRE(b == b); - REQUIRE(b == c); - REQUIRE(b != d); - - REQUIRE(c != a); - REQUIRE(c == b); - REQUIRE(c == c); - REQUIRE(c != d); - - REQUIRE(d != a); - REQUIRE(d != b); - REQUIRE(d != c); - REQUIRE(d == d); - } - - SECTION("operator<<") { - fixed_vector a = {1, 2, 3, 4, 5}; - - std::stringstream stream; - stream << a; - REQUIRE(stream.str() == "{1, 2, 3, 4, 5}"); - } - - SECTION("concat") { - fixed_vector a = {1, 2, 3}; - fixed_vector b = {4, 5}; - fixed_vector c = concat(a, b); - - REQUIRE(c[0] == 1); - REQUIRE(c[1] == 2); - REQUIRE(c[2] == 3); - REQUIRE(c[3] == 4); - REQUIRE(c[4] == 5); - } -} \ No newline at end of file diff --git a/test/utils/test_function_ref.cpp b/test/utils/test_function_ref.cpp new file mode 100644 index 00000000..ea7a4ea2 --- /dev/null +++ b/test/utils/test_function_ref.cpp @@ -0,0 +1,32 @@ +#include "catch2/catch_all.hpp" + +#include "kmm/utils/function_ref.hpp" + +using namespace kmm; + +static int add_one(int x) { + return x + 1; +} + +TEST_CASE("FunctionRef with lambda") { + int counter = 0; + auto counting = [&counter](int x) mutable { + counter++; + return x + counter; + }; + + function_ref ref = counting; + CHECK(ref(1) == 2); + CHECK(ref(1) == 3); + CHECK(counter == 2); +} + +TEST_CASE("FunctionRef with function pointer") { + function_ref ref = add_one; + CHECK(ref(41) == 42); +} + +TEST_CASE("FunctionRef default state is null") { + function_ref ref; + CHECK(!ref); +} \ No newline at end of file diff --git a/test/utils/test_hash_utils.cpp b/test/utils/test_hash_utils.cpp deleted file mode 100644 index 5989d6ea..00000000 --- a/test/utils/test_hash_utils.cpp +++ /dev/null @@ -1,38 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/hash_utils.hpp" - -using namespace kmm; - -TEST_CASE("hash_combine") { - size_t seed = 0; - hash_combine(seed, int32_t(32)); - hash_combine(seed, double(32)); - hash_combine(seed, std::string("foo")); - hash_combine(seed, true); - REQUIRE(seed != 0); -} - -TEST_CASE("hash_combine_range") { - SECTION("items") { - size_t seed0 = 0; - std::array a {1, 2, 3}; - hash_combine_range(seed0, a.data(), a.data() + a.size()); - - size_t seed1 = 0; - hash_combine(seed1, 1); - hash_combine(seed1, 2); - hash_combine(seed1, 3); - - REQUIRE(seed0 == seed1); - } - - SECTION("empty range") { - size_t seed = 1337; - - std::array a; - hash_combine_range(seed, a.data(), a.data()); - - REQUIRE(seed == 1337); - } -} \ No newline at end of file diff --git a/test/utils/test_integer_fun.cpp b/test/utils/test_integer_fun.cpp deleted file mode 100644 index 2c134a05..00000000 --- a/test/utils/test_integer_fun.cpp +++ /dev/null @@ -1,90 +0,0 @@ -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/utils/integer_fun.hpp" - -using namespace kmm; - -using i32 = int32_t; -using u64 = uint64_t; - -TEST_CASE("div_floor") { - REQUIRE(div_floor(0, 5) == 0); - REQUIRE(div_floor(8, 4) == 2); - REQUIRE(div_floor(-8, 4) == -2); - REQUIRE(div_floor(8, -4) == -2); - REQUIRE(div_floor(-8, -4) == 2); - REQUIRE(div_floor(7, 4) == 1); - REQUIRE(div_floor(-7, 4) == -2); - REQUIRE(div_floor(7, -4) == -2); - REQUIRE(div_floor(-7, -4) == 1); - - REQUIRE(div_floor(0ULL, 5ULL) == 0ULL); - REQUIRE(div_floor(9ULL, 2ULL) == 4ULL); - REQUIRE(div_floor(15ULL, 3ULL) == 5ULL); - REQUIRE(div_floor(123'456'789ULL, 1'000ULL) == 123'456ULL); -} - -TEST_CASE("div_ceil") { - REQUIRE(div_ceil(0, 5) == 0); - REQUIRE(div_ceil(8, 4) == 2); - REQUIRE(div_ceil(-8, 4) == -2); - REQUIRE(div_ceil(8, -4) == -2); - REQUIRE(div_ceil(-8, -4) == 2); - REQUIRE(div_ceil(7, 4) == 2); - REQUIRE(div_ceil(-7, 4) == -1); - REQUIRE(div_ceil(7, -4) == -1); - REQUIRE(div_ceil(-7, -4) == 2); - - REQUIRE(div_ceil(0ULL, 5ULL) == 0ULL); - REQUIRE(div_ceil(9ULL, 2ULL) == 5ULL); - REQUIRE(div_ceil(15ULL, 3ULL) == 5ULL); - REQUIRE(div_ceil(123'456'789ULL, 1'000ULL) == 123'457ULL); -} - -TEST_CASE("round_up_to_multiple") { - // positive input - REQUIRE(round_up_to_multiple(0, 8) == 0); - REQUIRE(round_up_to_multiple(7, 4) == 8); - REQUIRE(round_up_to_multiple(12, 4) == 12); - - // negative input - REQUIRE(round_up_to_multiple(-7, 4) == -4); - REQUIRE(round_up_to_multiple(-12, 4) == -12); - REQUIRE(round_up_to_multiple(-7, -4) == -4); - REQUIRE(round_up_to_multiple(-12, -4) == -12); - - // unsigned - REQUIRE(round_up_to_multiple(9ULL, 8ULL) == 16ULL); - REQUIRE(round_up_to_multiple(32ULL, 8ULL) == 32ULL); -} - -TEST_CASE("round_up_to_power_of_two") { - REQUIRE(round_up_to_power_of_two(-1) == 1); - REQUIRE(round_up_to_power_of_two(0) == 1); - REQUIRE(round_up_to_power_of_two(1) == 1); - REQUIRE(round_up_to_power_of_two(2) == 2); - REQUIRE(round_up_to_power_of_two(3) == 4); - REQUIRE(round_up_to_power_of_two(17) == 32); - - REQUIRE(round_up_to_power_of_two(63ULL) == 64ULL); - REQUIRE(round_up_to_power_of_two(64ULL) == 64ULL); - REQUIRE(round_up_to_power_of_two(65ULL) == 128ULL); -} - -TEST_CASE("is_power_of_two") { - REQUIRE(is_power_of_two(-1) == false); - REQUIRE(is_power_of_two(0) == false); - REQUIRE(is_power_of_two(1) == true); - REQUIRE(is_power_of_two(2) == true); - REQUIRE(is_power_of_two(3) == false); - REQUIRE(is_power_of_two(16) == true); - REQUIRE(is_power_of_two(std::numeric_limits::min()) == false); - REQUIRE(is_power_of_two(std::numeric_limits::max()) == false); - - REQUIRE(is_power_of_two(1'024ULL) == true); - REQUIRE(is_power_of_two(1'025ULL) == false); - REQUIRE(is_power_of_two(std::numeric_limits::min()) == false); - REQUIRE(is_power_of_two(std::numeric_limits::max()) == false); -} \ No newline at end of file diff --git a/test/utils/test_intrusive_ptr.cpp b/test/utils/test_intrusive_ptr.cpp new file mode 100644 index 00000000..1f4c23bc --- /dev/null +++ b/test/utils/test_intrusive_ptr.cpp @@ -0,0 +1,5 @@ +#include +#include +#include + +#include "catch2/catch_all.hpp" diff --git a/test/utils/test_key_value.cpp b/test/utils/test_key_value.cpp deleted file mode 100644 index cbd83188..00000000 --- a/test/utils/test_key_value.cpp +++ /dev/null @@ -1,96 +0,0 @@ -#include -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/utils/key_value.hpp" - -using namespace kmm; - -TEST_CASE("KeyValue") { - SECTION("int") { - KeyValue a {1, 123}; - KeyValue b {2, 123}; - KeyValue c {2, 456}; - - REQUIRE(a.key == 1); - REQUIRE(b.key == 2); - REQUIRE(c.key == 2); - REQUIRE(a.value == 123); - REQUIRE(b.value == 123); - REQUIRE(c.value == 456); - - REQUIRE(a == a); - REQUIRE(a != b); - REQUIRE(a != c); - REQUIRE(b != a); - REQUIRE(b == b); - REQUIRE(b != c); - REQUIRE(c != a); - REQUIRE(c != b); - REQUIRE(c == c); - - REQUIRE_FALSE(a < a); - REQUIRE(a < b); - REQUIRE(a < c); - REQUIRE_FALSE(b < a); - REQUIRE_FALSE(b < b); - REQUIRE(b < c); - REQUIRE_FALSE(c < a); - REQUIRE_FALSE(c < b); - REQUIRE_FALSE(c < c); - - REQUIRE(a <= a); - REQUIRE(a <= b); - REQUIRE(a <= c); - REQUIRE_FALSE(b <= a); - REQUIRE(b <= b); - REQUIRE(b <= c); - REQUIRE_FALSE(c <= a); - REQUIRE_FALSE(c <= b); - REQUIRE(c <= c); - } - - SECTION("float") { - KeyValue a {1, 123.0f}; - KeyValue b {2, 123.0f}; - KeyValue c {3, std::numeric_limits::quiet_NaN()}; - - REQUIRE(a.key == 1); - REQUIRE(b.key == 2); - REQUIRE(c.key == 3); - REQUIRE(a.value == 123.0f); - REQUIRE(b.value == 123.0f); - REQUIRE(std::isnan(c.value)); - - REQUIRE(a == a); - REQUIRE(a != b); - REQUIRE(a != c); - REQUIRE(b != a); - REQUIRE(b == b); - REQUIRE_FALSE(b == c); - REQUIRE(c != a); - REQUIRE(c != b); - REQUIRE(c != c); // because of NaN - - REQUIRE_FALSE(a < a); - REQUIRE(a < b); - REQUIRE(a < c); - REQUIRE_FALSE(b < a); - REQUIRE_FALSE(b < b); - REQUIRE(b < c); - REQUIRE_FALSE(c < a); - REQUIRE_FALSE(c < b); - REQUIRE_FALSE(c < c); - - REQUIRE(a <= a); - REQUIRE(a <= b); - REQUIRE(a <= c); - REQUIRE_FALSE(b <= a); - REQUIRE(b <= b); - REQUIRE(b <= c); - REQUIRE_FALSE(c <= a); - REQUIRE_FALSE(c <= b); - REQUIRE(c <= c); - } -} diff --git a/test/utils/test_lru_cache.cpp b/test/utils/test_lru_cache.cpp new file mode 100644 index 00000000..d083f2c8 --- /dev/null +++ b/test/utils/test_lru_cache.cpp @@ -0,0 +1,104 @@ +#include + +#include "catch2/catch_all.hpp" + +#include "kmm/utils/lru_cache.hpp" + +using namespace kmm; + +TEST_CASE("lru_cache construction") { + lru_cache cache; + CHECK(cache.size() == 0); + CHECK(cache.is_empty()); +} + +TEST_CASE("lru_cache::insert and find") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(2, "two"); + + CHECK(cache.size() == 2); + CHECK(cache.contains(1)); + CHECK(cache.contains(2)); + + auto* value = cache.find(1); + REQUIRE(value != nullptr); + CHECK(*value == "one"); + + CHECK(cache.find(3) == nullptr); + CHECK_FALSE(cache.contains(3)); +} + +TEST_CASE("lru_cache::insert overwrites existing key") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(1, "uno"); + + CHECK(cache.size() == 1); + CHECK(*cache.find(1) == "uno"); +} + +TEST_CASE("lru_cache::find marks entry as most recently used") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(2, "two"); + + // Touch 1 via find so that 2 becomes the least recently used entry. + cache.find(1); + + CHECK(*cache.least_recently_used() == 2); +} + +TEST_CASE("lru_cache::touch marks entry as most recently used") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(2, "two"); + CHECK(*cache.least_recently_used() == 1); + + cache.touch(1); + CHECK(*cache.least_recently_used() == 2); + + // touching a missing key is a no-op + cache.touch(99); + CHECK(*cache.least_recently_used() == 2); +} + +TEST_CASE("lru_cache::least_recently_used") { + lru_cache cache; + CHECK(cache.least_recently_used() == nullptr); + + cache.insert(1, "one"); + cache.insert(2, "two"); + + REQUIRE(cache.least_recently_used() != nullptr); + CHECK(*cache.least_recently_used() == 1); +} + +TEST_CASE("lru_cache::remove") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(2, "two"); + + CHECK(cache.remove(1)); + CHECK_FALSE(cache.contains(1)); + CHECK(cache.size() == 1); + + CHECK_FALSE(cache.remove(1)); +} + +TEST_CASE("lru_cache::clear") { + lru_cache cache; + + cache.insert(1, "one"); + cache.insert(2, "two"); + cache.clear(); + + CHECK(cache.is_empty()); + CHECK(cache.size() == 0); + CHECK(cache.least_recently_used() == nullptr); +} diff --git a/test/utils/test_panic.cpp b/test/utils/test_panic.cpp deleted file mode 100644 index 32c46e16..00000000 --- a/test/utils/test_panic.cpp +++ /dev/null @@ -1,8 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/panic.hpp" - -// Unfortunately, Catch2 does not allow catching SIGABRT. Maybe in the future... -TEST_CASE("panic", "[.shouldskip]") { - KMM_PANIC("don't worry, this is just a test"); -} \ No newline at end of file diff --git a/test/utils/test_point.cpp b/test/utils/test_point.cpp deleted file mode 100644 index f9bc82d1..00000000 --- a/test/utils/test_point.cpp +++ /dev/null @@ -1,128 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/utils/point.hpp" - -using namespace kmm; - -TEST_CASE("Point") { - SECTION("constructor") { - Point<3> a; - Point<3> b = {-1, 1, 2}; - - REQUIRE(a[0] == 0); - REQUIRE(a[1] == 0); - REQUIRE(a[2] == 0); - - REQUIRE(b[0] == -1); - REQUIRE(b[1] == 1); - REQUIRE(b[2] == 2); - } - - SECTION("conversions") { - Point<3> a = {1, 2, 3}; - Point<3> b = {-1, 1, 0}; - - REQUIRE_FALSE(a.is_convertible_to<2>()); - REQUIRE(a.is_convertible_to<4>()); - REQUIRE_FALSE(a.is_convertible_to<2, unsigned int>()); - REQUIRE(a.is_convertible_to<3, unsigned int>()); - REQUIRE(a.is_convertible_to<4, unsigned int>()); - - REQUIRE(b.is_convertible_to<2>()); - REQUIRE(b.is_convertible_to<4>()); - REQUIRE_FALSE(b.is_convertible_to<2, unsigned int>()); - REQUIRE_FALSE(b.is_convertible_to<3, unsigned int>()); - REQUIRE_FALSE(b.is_convertible_to<4, unsigned int>()); - - REQUIRE_THROWS(Point<2> {a}); - REQUIRE(Point<4> {a} == Point<4> {1, 2, 3, 0}); - REQUIRE_THROWS(Point<2, unsigned int> {a}); - REQUIRE(Point<4, unsigned int> {a} == Point<4, unsigned int> {1, 2, 3, 0}); - - REQUIRE(Point<2> {b} == Point<2> {-1, 1}); - REQUIRE(Point<4> {b} == Point<4> {-1, 1, 0, 0}); - REQUIRE_THROWS(Point<2, unsigned int> {b}); - REQUIRE_THROWS(Point<4, unsigned int> {b}); - - REQUIRE(Point<2>::from(a) == Point<2> {1, 2}); - REQUIRE(Point<4>::from(a) == Point<4> {1, 2, 3, 0}); - REQUIRE(Point<2, unsigned int>::from(a) == Point<2, unsigned int> {1, 2}); - REQUIRE(Point<4, unsigned int>::from(a) == Point<4, unsigned int> {1, 2, 3, 0}); - - REQUIRE(Point<2>::from(b) == Point<2> {-1, 1}); - REQUIRE(Point<4>::from(b) == Point<4> {-1, 1, 0, 0}); - REQUIRE(Point<2, unsigned int>::from(b) == Point<2, unsigned int> {UINT_MAX, 1}); - REQUIRE(Point<4, unsigned int>::from(b) == Point<4, unsigned int> {UINT_MAX, 1, 0, 0}); - } - - SECTION("fill") { - auto fill = Point<3>::fill(1337); - auto one = Point<3>::one(); - auto zero = Point<3>::zero(); - - REQUIRE(fill == Point<3> {1337, 1337, 1337}); - REQUIRE(one == Point<3> {1, 1, 1}); - REQUIRE(zero == Point<3> {0, 0, 0}); - } - - SECTION("get_or_default") { - Point<3> a; - Point<3> b = {-1, 1, 2}; - - REQUIRE(a.get_or_default(0) == 0); - REQUIRE(a.get_or_default(1) == 0); - REQUIRE(a.get_or_default(2) == 0); - REQUIRE(a.get_or_default(3) == 0); - REQUIRE(a.get_or_default(3, 1337) == 1337); - - REQUIRE(b.get_or_default(0) == -1); - REQUIRE(b.get_or_default(1) == 1); - REQUIRE(b.get_or_default(2) == 2); - REQUIRE(b.get_or_default(3) == 0); - REQUIRE(b.get_or_default(3, 1337) == 1337); - } - - SECTION("operator==") { - Point<3> a = {-1, 2, 3}; - Point<2> b = {-1, 2}; - Point<3> c = {-1, 2, 0}; - Point<3, uint64_t> d = {uint64_t(-1), 2, 3}; - Point<3, double> e = {-1, 2, 3}; - - REQUIRE(a == a); - REQUIRE(a != b); - REQUIRE(a != c); - REQUIRE(a != d); - REQUIRE(a == e); - - REQUIRE(b != a); - REQUIRE(b == b); - REQUIRE(b == c); - REQUIRE(b != d); - REQUIRE(b != e); - - REQUIRE(c != a); - REQUIRE(c == b); - REQUIRE(c == c); - REQUIRE(c != d); - REQUIRE(c != e); - - REQUIRE(d != a); - REQUIRE(d != b); - REQUIRE(d != c); - REQUIRE(d == d); - REQUIRE(d != e); - - REQUIRE(e == a); - REQUIRE(e != b); - REQUIRE(e != c); - REQUIRE(e != d); - REQUIRE(e == e); - } - - SECTION("operator<<") { - std::stringstream stream; - stream << Point {1, 2, 3}; - REQUIRE(stream.str() == "{1, 2, 3}"); - } -} \ No newline at end of file diff --git a/test/utils/test_range.cpp b/test/utils/test_range.cpp deleted file mode 100644 index b1458232..00000000 --- a/test/utils/test_range.cpp +++ /dev/null @@ -1,139 +0,0 @@ -#include - -#include "catch2/catch_all.hpp" - -#include "kmm/utils/range.hpp" - -using namespace kmm; - -TEST_CASE("range") { - Range empty; - REQUIRE(empty.begin == 0); - REQUIRE(empty.end == 0); - - Range one = {8}; - REQUIRE(one.begin == 0); - REQUIRE(one.end == 8); - - Range middle = {5, 10}; - REQUIRE(middle.begin == 5); - REQUIRE(middle.end == 10); - - Range max = {INT_MIN, INT_MAX}; - REQUIRE(max.begin == INT_MIN); - REQUIRE(max.end == INT_MAX); - - REQUIRE(empty.is_empty()); - REQUIRE_FALSE(one.is_empty()); - REQUIRE_FALSE(middle.is_empty()); - REQUIRE_FALSE(max.is_empty()); - - REQUIRE(empty.is_convertible_to()); - REQUIRE(one.is_convertible_to()); - REQUIRE(middle.is_convertible_to()); - REQUIRE_FALSE(max.is_convertible_to()); - - REQUIRE(Range(empty) == empty); - REQUIRE(Range(one) == one); - REQUIRE(Range(middle) == middle); - REQUIRE_THROWS(Range(max)); - - REQUIRE(Range::from(empty) == empty); - REQUIRE(Range::from(one) == one); - REQUIRE(Range::from(middle) == middle); - REQUIRE(Range::from(max) == Range(UINT_MAX / 2 + 1, UINT_MAX / 2)); - - REQUIRE_FALSE(empty.contains(-1)); - REQUIRE_FALSE(empty.contains(0)); - REQUIRE_FALSE(empty.contains(1)); - REQUIRE_FALSE(empty.contains(-100)); - REQUIRE_FALSE(one.contains(-1)); - REQUIRE(one.contains(0)); - REQUIRE(one.contains(1)); - REQUIRE_FALSE(one.contains(-100)); - REQUIRE_FALSE(middle.contains(-1)); - REQUIRE_FALSE(middle.contains(0)); - REQUIRE_FALSE(middle.contains(1)); - REQUIRE_FALSE(middle.contains(-100)); - REQUIRE(max.contains(-1)); - REQUIRE(max.contains(0)); - REQUIRE(max.contains(1)); - REQUIRE(max.contains(-100)); - REQUIRE(max.contains(INT_MIN)); - REQUIRE_FALSE(max.contains(INT_MAX)); - - REQUIRE(empty.contains(empty)); - REQUIRE_FALSE(empty.contains(one)); - REQUIRE_FALSE(empty.contains(middle)); - REQUIRE_FALSE(empty.contains(max)); - REQUIRE(one.contains(empty)); - REQUIRE(one.contains(one)); - REQUIRE_FALSE(one.contains(middle)); - REQUIRE_FALSE(one.contains(max)); - REQUIRE(middle.contains(empty)); - REQUIRE_FALSE(middle.contains(one)); - REQUIRE(middle.contains(middle)); - REQUIRE_FALSE(middle.contains(max)); - REQUIRE(max.contains(empty)); - REQUIRE(max.contains(one)); - REQUIRE(max.contains(middle)); - REQUIRE(max.contains(max)); - - REQUIRE_FALSE(empty.overlaps(empty)); - REQUIRE_FALSE(empty.overlaps(one)); - REQUIRE_FALSE(empty.overlaps(middle)); - REQUIRE_FALSE(empty.overlaps(max)); - REQUIRE_FALSE(one.overlaps(empty)); - REQUIRE(one.overlaps(one)); - REQUIRE(one.overlaps(middle)); - REQUIRE(one.overlaps(max)); - REQUIRE_FALSE(middle.overlaps(empty)); - REQUIRE(middle.overlaps(one)); - REQUIRE(middle.overlaps(middle)); - REQUIRE(middle.overlaps(max)); - REQUIRE_FALSE(max.overlaps(empty)); - REQUIRE(max.overlaps(one)); - REQUIRE(max.overlaps(middle)); - REQUIRE(max.overlaps(max)); - - REQUIRE(empty.intersection(empty) == empty); - REQUIRE(empty.intersection(one) == empty); - REQUIRE(empty.intersection(middle) == Range {5, 0}); - REQUIRE(empty.intersection(max) == empty); - REQUIRE(one.intersection(empty) == empty); - REQUIRE(one.intersection(one) == one); - REQUIRE(one.intersection(middle) == Range {5, 8}); - REQUIRE(one.intersection(max) == one); - REQUIRE(middle.intersection(empty) == Range {5, 0}); - REQUIRE(middle.intersection(one) == Range {5, 8}); - REQUIRE(middle.intersection(middle) == middle); - REQUIRE(middle.intersection(max) == middle); - REQUIRE(max.intersection(empty) == empty); - REQUIRE(max.intersection(one) == one); - REQUIRE(max.intersection(middle) == middle); - REQUIRE(max.intersection(max) == max); - - REQUIRE(empty.size() == 0); - REQUIRE(one.size() == 8); - REQUIRE(middle.size() == 5); - // REQUIRE(max.size() == -1); // ?? - - SECTION("split_tail") { - auto before = Range {5, 25}; - auto after = before.split_tail(10); - REQUIRE(before == Range {5, 10}); - REQUIRE(after == Range {10, 25}); - } - - SECTION("shift_by") { - auto total = Range {5, 25}; - auto shifted = total.shift_by(10); - REQUIRE(shifted == Range {15, 35}); - } - - SECTION("operator<<") { - std::stringstream ss; - ss << middle; - REQUIRE(ss.str() == "5...10"); - } -} \ No newline at end of file diff --git a/test/utils/test_small_vector.cpp b/test/utils/test_small_vector.cpp index 136995b3..5f3d03f5 100644 --- a/test/utils/test_small_vector.cpp +++ b/test/utils/test_small_vector.cpp @@ -1,5 +1,5 @@ +#include #include -#include #include "catch2/catch_all.hpp" @@ -7,166 +7,235 @@ using namespace kmm; -struct MyString { - MyString(std::string x = "") : value(x) {} - std::string value; -}; - -TEST_CASE("small_vector") { - SECTION("basics") { - small_vector x; - REQUIRE(x.capacity() == 4); - REQUIRE(x.size() == 0); - REQUIRE(x.is_empty() == true); - REQUIRE(x.is_heap_allocated() == false); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data()); - - x.push_back(1); - - REQUIRE(x.capacity() == 4); - REQUIRE(x.size() == 1); - REQUIRE(x.is_empty() == false); - REQUIRE(x.is_heap_allocated() == false); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data() + 1); - REQUIRE(&x[0] == x.data()); - REQUIRE(x[0] == 1); - - x.push_back(2); - x.push_back(3); - x.push_back(4); - x.push_back(5); - - REQUIRE(x.capacity() == 16); - REQUIRE(x.size() == 5); - REQUIRE(x.is_empty() == false); - REQUIRE(x.is_heap_allocated() == true); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data() + 5); - REQUIRE(&x[0] == x.data()); - REQUIRE(x[0] == 1); - REQUIRE(x[1] == 2); - REQUIRE(x[2] == 3); - REQUIRE(x[3] == 4); - REQUIRE(x[4] == 5); - - x.truncate(2); - - REQUIRE(x.capacity() == 16); - REQUIRE(x.size() == 2); - REQUIRE(x.is_empty() == false); - REQUIRE(x.is_heap_allocated() == true); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data() + 2); - REQUIRE(&x[0] == x.data()); - REQUIRE(x[0] == 1); - REQUIRE(x[1] == 2); - - x.resize(20); - - REQUIRE(x.capacity() == 32); - REQUIRE(x.size() == 20); - REQUIRE(x.is_empty() == false); - REQUIRE(x.is_heap_allocated() == true); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data() + 20); - REQUIRE(&x[0] == x.data()); - REQUIRE(x[0] == 1); - REQUIRE(x[1] == 2); - REQUIRE(x[2] == 0); - REQUIRE(x[3] == 0); - REQUIRE(x[4] == 0); - REQUIRE(x[5] == 0); - REQUIRE(x[6] == 0); - REQUIRE(x[7] == 0); - REQUIRE(x[19] == 0); - - x.clear(); - - REQUIRE(x.capacity() == 32); - REQUIRE(x.size() == 0); - REQUIRE(x.is_empty() == true); - REQUIRE(x.is_heap_allocated() == true); - REQUIRE(x.begin() == x.data()); - REQUIRE(x.end() == x.data()); +TEST_CASE("small_vector construction") { + small_vector a; + CHECK(a.size() == 0); + CHECK(a.is_empty()); + CHECK(a.capacity() == 4); + CHECK_FALSE(a.is_heap_allocated()); + + small_vector b {1, 2, 3}; + CHECK(b.size() == 3); + CHECK_FALSE(b.is_empty()); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); +} + +TEST_CASE("small_vector::push_back") { + SECTION("within capacity") { + small_vector a; + + a.push_back(1); + a.push_back(2); + a.push_back(3); + + CHECK(a.size() == 3); + CHECK(a.capacity() == 4); + CHECK_FALSE(a.is_heap_allocated()); + CHECK(a[0] == 1); + CHECK(a[1] == 2); + CHECK(a[2] == 3); } - SECTION("constructor") { - small_vector a; // default - small_vector b = {"a", "b"}; // list - small_vector c = b; // copy - small_vector d = std::move(c); //move - small_vector e = d; // copy, different N - small_vector f = d; // copy, different T + SECTION("beyond capacity") { + small_vector a; - REQUIRE(a.size() == 0); + a.push_back(1); + a.push_back(2); + CHECK_FALSE(a.is_heap_allocated()); - REQUIRE(b.size() == 2); - REQUIRE(b[0] == "a"); - REQUIRE(b[1] == "b"); + a.push_back(3); + CHECK(a.is_heap_allocated()); + CHECK(a.size() == 3); + CHECK(a.capacity() >= 3); - REQUIRE(c.size() == 0); - - REQUIRE(d.size() == 2); - REQUIRE(d[0] == "a"); - REQUIRE(d[1] == "b"); - - REQUIRE(e.size() == 2); - REQUIRE(e[0] == "a"); - REQUIRE(e[1] == "b"); - - REQUIRE(f.size() == 2); - REQUIRE(f[0].value == "a"); - REQUIRE(f[1].value == "b"); - - // The other two values must be uninitialized - REQUIRE(f.capacity() == 4); - REQUIRE(f[2].value == ""); - REQUIRE(f[3].value == ""); + CHECK(a[0] == 1); + CHECK(a[1] == 2); + CHECK(a[2] == 3); + } +} + +TEST_CASE("small_vector::try_push_back") { + small_vector a; + + CHECK(a.try_push_back(1)); + CHECK(a.try_push_back(2)); + CHECK(a.try_push_back(3)); + + CHECK(a.size() == 3); + CHECK(a[2] == 3); +} + +TEST_CASE("small_vector copy constructor") { + small_vector a {1, 2, 3}; + small_vector b = a; + + CHECK(b.size() == 3); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); + + // modifying the copy should not affect the original + b[0] = 99; + CHECK(a[0] == 1); + CHECK(b[0] == 99); +} + +TEST_CASE("small_vector copy constructor with different data type and inline size") { + small_vector a {1, 2, 3}; + small_vector b = a; + + CHECK(b.size() == 3); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); +} + +TEST_CASE("small_vector copy assignment") { + small_vector a {1, 2, 3}; + small_vector b {9, 9}; + + b = a; + + CHECK(b.size() == 3); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); + + // self-assignment should be a no-op + b = b; + CHECK(b.size() == 3); + CHECK(b[0] == 1); +} + +TEST_CASE("small_vector move constructor") { + small_vector a {1, 2, 3}; + small_vector b = std::move(a); + + CHECK(b.size() == 3); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); +} + +TEST_CASE("small_vector move constructor with heap allocation") { + small_vector a {1, 2, 3, 4}; + CHECK(a.is_heap_allocated()); + + small_vector b = std::move(a); + + CHECK(b.is_heap_allocated()); + CHECK(b.size() == 4); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); + CHECK(b[3] == 4); +} + +TEST_CASE("small_vector move assignment") { + small_vector a {1, 2, 3}; + small_vector b {9}; + + b = std::move(a); + + CHECK(b.size() == 3); + CHECK(b[0] == 1); + CHECK(b[1] == 2); + CHECK(b[2] == 3); +} + +TEST_CASE("small_vector::resize") { + SECTION("within inline") { + small_vector a {1, 2}; + a.resize(4); + CHECK(a.size() == 4); + + a.resize(1); + CHECK(a.size() == 1); + CHECK(a[0] == 1); } - SECTION("operator=") { - small_vector a = {"foo", "bar"}; - small_vector b; - small_vector c; - small_vector d; - small_vector e; + SECTION("beyond inline") { + small_vector a {1, 2}; + a.resize(10); - b = a; // operator=(const small_vector&) - c = a; // operator=(const small_vector&) - d = a; // operator=(const small_vector&) - e = std::move(a); // operator=(small_vector&&) + CHECK(a.size() == 10); + CHECK(a.is_heap_allocated()); + CHECK(a[0] == 1); + CHECK(a[1] == 2); + } +} + +TEST_CASE("small_vector::truncate") { + small_vector a {1, 2, 3, 4}; + + a.truncate(2); + CHECK(a.size() == 2); + CHECK(a[0] == 1); + CHECK(a[1] == 2); + + // truncating to a larger size than current size should be a no-op + a.truncate(10); + CHECK(a.size() == 2); +} + +TEST_CASE("small_vector::clear") { + small_vector a {1, 2, 3}; + a.clear(); + + CHECK(a.size() == 0); + CHECK(a.is_empty()); + CHECK(a.capacity() == 4); +} + +TEST_CASE("small_vector::insert_all") { + SECTION("iterator") { + small_vector a {1, 2}; + int extra[] = {3, 4, 5}; + + a.insert_all(std::begin(extra), std::end(extra)); + + CHECK(a.size() == 5); + CHECK(a.is_heap_allocated()); + for (size_t i = 0; i < 5; i++) { + CHECK(a[i] == static_cast(i + 1)); + } + } - REQUIRE(a.size() == 0); + SECTION("from small_vector") { + small_vector a {1, 2}; + small_vector b {3, 4}; - REQUIRE(b.size() == 2); - REQUIRE(b[0] == "foo"); - REQUIRE(b[1] == "bar"); + a.insert_all(b); - REQUIRE(c.size() == 2); - REQUIRE(c[0] == "foo"); - REQUIRE(c[1] == "bar"); + CHECK(a.size() == 4); + CHECK(a[0] == 1); + CHECK(a[1] == 2); + CHECK(a[2] == 3); + CHECK(a[3] == 4); + } +} - REQUIRE(d.size() == 2); - REQUIRE(d[0].value == "foo"); - REQUIRE(d[1].value == "bar"); +TEST_CASE("small_vector iterator") { + small_vector a {1, 2, 3}; - REQUIRE(e.size() == 2); - REQUIRE(e[0] == "foo"); - REQUIRE(e[1] == "bar"); + int sum = 0; + for (int v : a) { + sum += v; } - SECTION("operator<<") { - small_vector x = {1, 2, 3, 4, 5, 6}; - small_vector y = {}; + CHECK(sum == 6); - auto stream = std::stringstream(); - stream << x; - REQUIRE(stream.str() == "{1, 2, 3, 4, 5, 6}"); + auto it = a.begin(); + CHECK(*it == 1); + CHECK(a.end() - a.begin() == 3); +} - stream = std::stringstream(); - stream << y; - REQUIRE(stream.str() == "{}"); - } +TEST_CASE("small_vector operator<<") { + small_vector a {1, 2, 3}; + CHECK(fmt::to_string(a) == "{1, 2, 3}"); + + small_vector empty; + CHECK(fmt::to_string(empty) == "{}"); } \ No newline at end of file diff --git a/test/utils/test_view.cpp b/test/utils/test_view.cpp deleted file mode 100644 index 6596241e..00000000 --- a/test/utils/test_view.cpp +++ /dev/null @@ -1,374 +0,0 @@ -#include "catch2/catch_all.hpp" - -#include "kmm/core/view.hpp" -#define CHECK_EQ(A, B) CHECK((A) == (B)) - -using namespace kmm; - -TEST_CASE("view, bound_left_to_right_layout") { - std::vector vec = {1, 2, 3, 4, 5, 6, 7, 8}; - AbstractView, views::left_to_right_layout<>> v = { - vec.data(), - {{8}} - }; - - CHECK_EQ(v.offset(), 0); - CHECK_EQ(v.size(0), 8); - CHECK_EQ(v.begin(), 0); - CHECK_EQ(v.end(), 8); - CHECK_EQ(v.data(), vec.data()); - CHECK_EQ(v.stride(), 1); - CHECK_EQ(v.strides(), 1); - CHECK_EQ(v.offsets(), 0); - CHECK_EQ(v.sizes(), 8); - - CHECK_EQ(v.data_at({0}), &vec[0]); - CHECK_EQ(v.data_at({4}), &vec[4]); - CHECK_EQ(v.data_at({8}), &vec[8]); - - CHECK_EQ(v.access({0}), vec[0]); - CHECK_EQ(v.access({4}), vec[4]); - - CHECK_EQ(v[0], vec[0]); - CHECK_EQ(v[4], vec[4]); -} - -TEST_CASE("view, bound2_left_to_right_layout") { - std::vector vec = {1, 2, 3, 4, 5, 6, 7, 8}; - AbstractView, views::left_to_right_layout<>> v = { - vec.data(), - {{4, 2}} - }; - - CHECK_EQ(v.offset(0), 0); - CHECK_EQ(v.offset(1), 0); - CHECK_EQ(v.size(0), 4); - CHECK_EQ(v.size(1), 2); - CHECK_EQ(v.begin(0), 0); - CHECK_EQ(v.begin(1), 0); - CHECK_EQ(v.end(0), 4); - CHECK_EQ(v.end(1), 2); - CHECK_EQ(v.data(), vec.data()); - CHECK_EQ(v.stride(0), 1); - CHECK_EQ(v.stride(1), 4); - CHECK_EQ(v.strides()[0], 1); - CHECK_EQ(v.strides()[1], 4); - CHECK_EQ(v.offsets()[0], 0); - CHECK_EQ(v.offsets()[1], 0); - CHECK_EQ(v.sizes()[0], 4); - CHECK_EQ(v.sizes()[1], 2); - - CHECK_EQ(v.data_at({0, 0}), &vec[0]); - CHECK_EQ(v.data_at({1, 1}), &vec[5]); - CHECK_EQ(v.data_at({3, 1}), &vec[7]); - CHECK_EQ(v.data_at({3, 2}), &vec[11]); - - CHECK_EQ(v.access({0, 1}), vec[4]); - CHECK_EQ(v.access({3, 1}), vec[7]); - - CHECK_EQ(v[0][1], vec[4]); - CHECK_EQ(v[3][0], vec[3]); -} - -TEST_CASE("view, bound2_right_to_left_layout") { - std::vector vec = {1, 2, 3, 4, 5, 6, 7, 8}; - AbstractView, views::right_to_left_layout<>> v = { - vec.data(), - {{4, 2}} - }; - - CHECK_EQ(v.offset(0), 0); - CHECK_EQ(v.offset(1), 0); - CHECK_EQ(v.size(0), 4); - CHECK_EQ(v.size(1), 2); - CHECK_EQ(v.begin(0), 0); - CHECK_EQ(v.begin(1), 0); - CHECK_EQ(v.end(0), 4); - CHECK_EQ(v.end(1), 2); - CHECK_EQ(v.data(), vec.data()); - CHECK_EQ(v.stride(0), 2); - CHECK_EQ(v.stride(1), 1); - CHECK_EQ(v.strides()[0], 2); - CHECK_EQ(v.strides()[1], 1); - CHECK_EQ(v.offsets()[0], 0); - CHECK_EQ(v.offsets()[1], 0); - CHECK_EQ(v.sizes()[0], 4); - CHECK_EQ(v.sizes()[1], 2); - - CHECK_EQ(v.data_at({0, 0}), &vec[0]); - CHECK_EQ(v.data_at({1, 1}), &vec[3]); - CHECK_EQ(v.data_at({3, 1}), &vec[7]); - CHECK_EQ(v.data_at({3, 2}), &vec[8]); - - CHECK_EQ(v.access({0, 1}), vec[1]); - CHECK_EQ(v.access({3, 1}), vec[7]); - - CHECK_EQ(v[0][1], vec[1]); - CHECK_EQ(v[3][0], vec[6]); -} - -TEST_CASE("view, subbound2_right_to_left_layout") { - std::vector vec = {1, 2, 3, 4, 5, 6, 7, 8}; - AbstractView, views::right_to_left_layout<>> v = { - vec.data(), - {{100, 42}, {4, 2}} - }; - - CHECK_EQ(v.offset(0), 100); - CHECK_EQ(v.offset(1), 42); - CHECK_EQ(v.size(0), 4); - CHECK_EQ(v.size(1), 2); - CHECK_EQ(v.begin(0), 100); - CHECK_EQ(v.begin(1), 42); - CHECK_EQ(v.end(0), 104); - CHECK_EQ(v.end(1), 44); - CHECK_EQ(v.data(), vec.data()); - CHECK_EQ(v.stride(0), 2); - CHECK_EQ(v.stride(1), 1); - CHECK_EQ(v.strides()[0], 2); - CHECK_EQ(v.strides()[1], 1); - CHECK_EQ(v.offsets()[0], 100); - CHECK_EQ(v.offsets()[1], 42); - CHECK_EQ(v.sizes()[0], 4); - CHECK_EQ(v.sizes()[1], 2); - - CHECK_EQ(v.data_at({100, 42}), &vec[0]); - CHECK_EQ(v.data_at({101, 43}), &vec[3]); - CHECK_EQ(v.data_at({103, 43}), &vec[7]); - CHECK_EQ(v.data_at({103, 44}), &vec[8]); - - CHECK_EQ(v.access({100, 43}), vec[1]); - CHECK_EQ(v.access({103, 43}), vec[7]); - - CHECK_EQ(v[100][43], vec[1]); - CHECK_EQ(v[103][42], vec[6]); -} - -TEST_CASE("view, domain_conversions") { -#define CHECK_CORRECT_VIEW(p) \ - CHECK_EQ((p).offset(0), 0); \ - CHECK_EQ((p).offset(1), 0); \ - CHECK_EQ((p).size(0), 10); \ - CHECK_EQ((p).size(1), 20); \ - CHECK_EQ((p).stride(0), 20); \ - CHECK_EQ((p).stride(1), 1); \ - CHECK((p).is_contiguous()); - - auto a = AbstractView< // - int, - views::static_domain, - views::right_to_left_layout<>> {nullptr}; - CHECK_CORRECT_VIEW(a); - - auto b = AbstractView< // - int, - views::dynamic_domain<2>, - views::right_to_left_layout<>>(a); - CHECK_CORRECT_VIEW(b); - - auto c = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::right_to_left_layout<>>(a); - CHECK_CORRECT_VIEW(c); - - auto d = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::right_to_left_layout<>>(b); - CHECK_CORRECT_VIEW(d); - - auto e = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::strided_layout<>>(a); - CHECK_CORRECT_VIEW(e); - - auto f = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::strided_layout<>>(b); - CHECK_CORRECT_VIEW(f); - - auto g = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::strided_layout<>>(c); - CHECK_CORRECT_VIEW(g); - - auto h = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::strided_layout<>>(d); - CHECK_CORRECT_VIEW(h); - -#undef CHECK_CORRECT_VIEW -} - -TEST_CASE("view, subdomain_conversions") { -#define CHECK_CORRECT_VIEW(p) \ - CHECK_EQ((p).offset(0), 3); \ - CHECK_EQ((p).offset(1), 7); \ - CHECK_EQ((p).size(0), 10); \ - CHECK_EQ((p).size(1), 20); \ - CHECK_EQ((p).stride(0), 20); \ - CHECK_EQ((p).stride(1), 1); \ - CHECK((p).is_contiguous()); - - auto a = AbstractView< // - int, - views::static_offset, 3, 7>, - views::right_to_left_layout<>> {nullptr}; - CHECK_CORRECT_VIEW(a); - - auto b = AbstractView< // - int, - views::static_offset, 3, 7>, - views::right_to_left_layout<>>(a); - CHECK_CORRECT_VIEW(b); - - auto c = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::right_to_left_layout<>>(a); - CHECK_CORRECT_VIEW(c); - - auto d = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::right_to_left_layout<>>(b); - CHECK_CORRECT_VIEW(d); - - auto e = AbstractView< // - int, - views::dynamic_subdomain<2>, - views::strided_layout<>>(a); - CHECK_CORRECT_VIEW(e); - - auto f = AbstractView, views::strided_layout<>>(b); - CHECK_CORRECT_VIEW(f); - - auto g = AbstractView, views::strided_layout<>>(c); - CHECK_CORRECT_VIEW(g); - - auto h = AbstractView, views::strided_layout<>>(d); - CHECK_CORRECT_VIEW(h); - -#undef CHECK_CORRECT_VIEW -} - -TEST_CASE("view, drop_axis_dim2") { - auto vec = std::vector(200); - auto a = AbstractView< // - float, - views::dynamic_subdomain<2>, - views::right_to_left_layout<>> {vec.data(), {{3, 7}, {10, 20}}}; - - // Drop axis 0 - AbstractView, views::right_to_left_layout<>> b = - a.drop_axis(); - CHECK_EQ(b.size(0), 20); - CHECK_EQ(b.offset(0), 7); - CHECK_EQ(b.stride(0), 1); - CHECK_EQ(b.data(), vec.data()); - - b = a.drop_axis(5); - CHECK_EQ(b.size(0), 20); - CHECK_EQ(b.offset(0), 7); - CHECK_EQ(b.stride(0), 1); - CHECK_EQ(b.data(), vec.data() + 2 * 20); - - // Drop axis 1 - AbstractView, views::strided_layout<>> c = a.drop_axis<1>(); - CHECK_EQ(c.size(0), 10); - CHECK_EQ(c.offset(0), 3); - CHECK_EQ(c.stride(0), 20); - CHECK_EQ(c.data(), vec.data()); - - c = a.drop_axis<1>(13); - CHECK_EQ(c.size(0), 10); - CHECK_EQ(c.offset(0), 3); - CHECK_EQ(c.stride(0), 20); - CHECK_EQ(c.data(), vec.data() + 6); -} - -TEST_CASE("view, drop_axis_dim3") { - auto vec = std::vector(200); - auto a = AbstractView< // - float, - views::dynamic_subdomain<3>, - views::right_to_left_layout<>> {vec.data(), {{3, 7, 1}, {2, 5, 20}}}; - - // Drop axis 0 - AbstractView, views::right_to_left_layout<>> b = - a.drop_axis(); - CHECK_EQ(b.size(0), 5); - CHECK_EQ(b.offset(0), 7); - CHECK_EQ(b.stride(0), 20); - CHECK_EQ(b.size(1), 20); - CHECK_EQ(b.offset(1), 1); - CHECK_EQ(b.stride(1), 1); - CHECK_EQ(b.data(), vec.data()); - - b = a.drop_axis(4); - CHECK_EQ(b.size(0), 5); - CHECK_EQ(b.offset(0), 7); - CHECK_EQ(b.stride(0), 20); - CHECK_EQ(b.size(1), 20); - CHECK_EQ(b.offset(1), 1); - CHECK_EQ(b.stride(1), 1); - CHECK_EQ(b.data() - vec.data(), 100); - - // Drop axis 1 - AbstractView, views::strided_layout<>> c = a.drop_axis<1>(); - CHECK_EQ(c.size(0), 2); - CHECK_EQ(c.offset(0), 3); - CHECK_EQ(c.stride(0), 100); - CHECK_EQ(c.size(1), 20); - CHECK_EQ(c.offset(1), 1); - CHECK_EQ(c.stride(1), 1); - CHECK_EQ(c.data(), vec.data()); - - c = a.drop_axis<1>(9); - CHECK_EQ(c.size(0), 2); - CHECK_EQ(c.offset(0), 3); - CHECK_EQ(c.stride(0), 100); - CHECK_EQ(c.size(1), 20); - CHECK_EQ(c.offset(1), 1); - CHECK_EQ(c.stride(1), 1); - CHECK_EQ(c.data() - vec.data(), +40); - - // Drop axis 3 - AbstractView, views::strided_layout<>> d = a.drop_axis<2>(); - CHECK_EQ(d.size(0), 2); - CHECK_EQ(d.offset(0), 3); - CHECK_EQ(d.stride(0), 100); - CHECK_EQ(d.size(1), 5); - CHECK_EQ(d.offset(1), 7); - CHECK_EQ(d.stride(1), 20); - CHECK_EQ(d.data(), vec.data()); - - d = a.drop_axis<2>(13); - CHECK_EQ(d.size(0), 2); - CHECK_EQ(d.offset(0), 3); - CHECK_EQ(d.stride(0), 100); - CHECK_EQ(d.size(1), 5); - CHECK_EQ(d.offset(1), 7); - CHECK_EQ(d.stride(1), 20); - CHECK_EQ(d.data() - vec.data(), 12); -} - -TEST_CASE("view, scalar") { - auto value = int(1); - auto v = AbstractView< // - int, - views::dynamic_subdomain<0>, - views::right_to_left_layout<>> {&value}; - - CHECK_EQ(v.data(), &value); - CHECK_EQ(v.data_at({}), &value); - CHECK_EQ(v.access({}), value); - - *v = 2; - CHECK_EQ(value, 2); -} \ No newline at end of file