//===----------------------------------------------------------------------===// // // Part of CUDA Experimental in CUDA C++ Core Libraries, // under the Apache License v2.0 with LLVM Exceptions. // See https://llvm.org/LICENSE.txt for license information. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception // SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. // //===----------------------------------------------------------------------===// #pragma once #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include inline std::string inspect_sass(const void* cubin, size_t cubin_size) { namespace fs = std::filesystem; fs::path temp_dir = fs::temp_directory_path(); fs::path temp_in_filename = temp_dir / "temp_in_file.cubin"; fs::path temp_out_filename = temp_dir / "temp_out_file.sass"; std::ofstream temp_in_file(temp_in_filename, std::ios::binary); if (!temp_in_file) { throw std::runtime_error("Failed to create temporary file."); } temp_in_file.write(static_cast(cubin), static_cast(cubin_size)); temp_in_file.close(); std::string command = "nvdisasm -gi "; command += temp_in_filename.string(); command += " > "; command += temp_out_filename.string(); int exec_code = std::system(command.c_str()); if (!fs::remove(temp_in_filename)) { throw std::runtime_error("Failed to remove temporary file."); } if (exec_code != 0) { throw std::runtime_error("Failed to execute command."); } std::ifstream temp_out_file(temp_out_filename, std::ios::binary); if (!temp_out_file) { throw std::runtime_error("Failed to create temporary file."); } const std::string sass{std::istreambuf_iterator(temp_out_file), std::istreambuf_iterator()}; if (!fs::remove(temp_out_filename)) { throw std::runtime_error("Failed to remove temporary file."); } return sass; } inline std::string compile(const std::string& source) { // Compile source to LTO-IR via NVRTC. v2 tests use the same path as v1; // the v2 backend then routes the LTO-IR through nvJitLink at link time. nvrtcProgram prog; REQUIRE(NVRTC_SUCCESS == nvrtcCreateProgram(&prog, source.c_str(), "op.cu", 0, nullptr, nullptr)); // TEST_CTK_PATH needed to include cuda_fp16.h const char* options[] = {"--std=c++17", "-rdc=true", "-dlto", "-D__NV_NO_VECTOR_DEPRECATION_DIAG", TEST_CTK_PATH}; if (nvrtcCompileProgram(prog, 5, options) != NVRTC_SUCCESS) { size_t log_size{}; REQUIRE(NVRTC_SUCCESS == nvrtcGetProgramLogSize(prog, &log_size)); std::vector log(log_size); REQUIRE(NVRTC_SUCCESS == nvrtcGetProgramLog(prog, log.data())); printf("%s\r\n", log.data()); REQUIRE(false); } std::size_t ltoir_size{}; REQUIRE(NVRTC_SUCCESS == nvrtcGetLTOIRSize(prog, <oir_size)); std::vector ltoir(ltoir_size); REQUIRE(NVRTC_SUCCESS == nvrtcGetLTOIR(prog, ltoir.data())); REQUIRE(NVRTC_SUCCESS == nvrtcDestroyProgram(&prog)); return std::string(ltoir.data(), ltoir_size); } // Helper to construct a cccl_build_config that works for both v1 (4-field) // and v2 (6-field) struct layouts. Uses value-init + explicit assignments so // v2's enable_pch / verbose get zero-initialized rather than left undefined. inline cccl_build_config make_build_config(const char** extra_flags, size_t num_flags, const char** extra_dirs, size_t num_dirs) { cccl_build_config config{}; config.extra_compile_flags = extra_flags; config.num_extra_compile_flags = num_flags; config.extra_include_dirs = extra_dirs; config.num_extra_include_dirs = num_dirs; return config; } template std::vector generate(std::size_t num_items) { // Add support for 8-bit ints, otherwise MSVC fails with: // error C2338: static_assert failed: // 'invalid template argument for uniform_int_distribution: // N4950 [rand.req.genl]/1.5 requires one of // short, int, long, long long, // unsigned short, unsigned int, unsigned long, or unsigned long long' using dist_type = std::conditional_t; std::random_device rnd_device; std::mt19937 mersenne_engine{rnd_device()}; // Generates random integers std::uniform_int_distribution dist{dist_type{1}, dist_type{42}}; std::vector vec(num_items); std::generate(vec.begin(), vec.end(), [&]() { return static_cast(dist(mersenne_engine)); }); return vec; } template std::vector make_shuffled_sequence(std::size_t num_items) { std::vector sequence(num_items); std::iota(sequence.begin(), sequence.end(), T(0)); std::random_device rnd_device; std::mt19937 mersenne_engine{rnd_device()}; std::shuffle(sequence.begin(), sequence.end(), mersenne_engine); return sequence; } template cccl_type_info get_type_info() { cccl_type_info info; info.size = sizeof(T); info.alignment = alignof(T); if constexpr (std::is_same_v || (std::is_integral_v && std::is_signed_v && sizeof(T) == sizeof(char))) { info.type = cccl_type_enum::CCCL_INT8; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_unsigned_v && sizeof(T) == sizeof(char) && !std::is_same_v) ) { info.type = cccl_type_enum::CCCL_UINT8; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_signed_v && sizeof(T) == sizeof(int16_t))) { info.type = cccl_type_enum::CCCL_INT16; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_unsigned_v && sizeof(T) == sizeof(int16_t))) { info.type = cccl_type_enum::CCCL_UINT16; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_signed_v && sizeof(T) == sizeof(int32_t))) { info.type = cccl_type_enum::CCCL_INT32; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_unsigned_v && sizeof(T) == sizeof(int32_t))) { info.type = cccl_type_enum::CCCL_UINT32; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_signed_v && sizeof(T) == sizeof(int64_t))) { info.type = cccl_type_enum::CCCL_INT64; } else if constexpr (std::is_same_v || (std::is_integral_v && std::is_unsigned_v && sizeof(T) == sizeof(int64_t))) { info.type = cccl_type_enum::CCCL_UINT64; } #if _CCCL_HAS_NVFP16() else if constexpr (std::is_same_v) { info.type = cccl_type_enum::CCCL_FLOAT16; } #endif else if constexpr (std::is_same_v) { info.type = cccl_type_enum::CCCL_FLOAT32; } else if constexpr (std::is_same_v) { info.type = cccl_type_enum::CCCL_FLOAT64; } else if constexpr (!std::is_integral_v) { info.type = cccl_type_enum::CCCL_STORAGE; } else { static_assert(false, "Unsupported type"); } return info; } std::string type_enum_to_name(cccl_type_enum type) { switch (type) { case cccl_type_enum::CCCL_INT8: return "char"; case cccl_type_enum::CCCL_INT16: return "short"; case cccl_type_enum::CCCL_INT32: return "int"; case cccl_type_enum::CCCL_INT64: return "long long"; case cccl_type_enum::CCCL_UINT8: return "unsigned char"; case cccl_type_enum::CCCL_UINT16: return "unsigned short"; case cccl_type_enum::CCCL_UINT32: return "unsigned int"; case cccl_type_enum::CCCL_UINT64: return "unsigned long long"; #if _CCCL_HAS_NVFP16() case cccl_type_enum::CCCL_FLOAT16: return "__half"; #endif case cccl_type_enum::CCCL_FLOAT32: return "float"; case cccl_type_enum::CCCL_FLOAT64: return "double"; default: throw std::runtime_error("Unsupported type"); } return ""; } // TOOD: using more than than one `op` in the same TU will fail because // of the lack of name mangling. Ditto for all `get_*_op` functions. inline std::string get_reduce_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " char* a = reinterpret_cast(a_void); " " char* b = reinterpret_cast(b_void); " " char* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " int* a = reinterpret_cast(a_void); " " int* b = reinterpret_cast(b_void); " " int* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " unsigned int* a = reinterpret_cast(a_void); " " unsigned int* b = reinterpret_cast(b_void); " " unsigned int* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " long long* a = reinterpret_cast(a_void); " " long long* b = reinterpret_cast(b_void); " " long long* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " unsigned long long* a = reinterpret_cast(a_void); " " unsigned long long* b = reinterpret_cast(b_void); " " unsigned long long* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_FLOAT32: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " float* a = reinterpret_cast(a_void); " " float* b = reinterpret_cast(b_void); " " float* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_FLOAT64: return "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " double* a = reinterpret_cast(a_void); " " double* b = reinterpret_cast(b_void); " " double* out = reinterpret_cast(out_void); " " *out = *a + *b; " "}"; case cccl_type_enum::CCCL_FLOAT16: return "#include \n" "extern \"C\" __device__ void op(void* a_void, void* b_void, void* out_void) { " " __half* a = reinterpret_cast<__half*>(a_void); " " __half* b = reinterpret_cast<__half*>(b_void); " " __half* out = reinterpret_cast<__half*>(out_void); " " *out = *a + *b; " "}"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::string get_for_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void op(void* a_void) { " " char* a = reinterpret_cast(a_void); " " (*a)++; " "}"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void op(void* a_void) { " " int* a = reinterpret_cast(a_void); " " (*a)++; " "}"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void op(void* a_void) { " " unsigned int* a = reinterpret_cast(a_void); " " (*a)++; " "}"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void op(void* a_void) { " " long long* a = reinterpret_cast(a_void); " " (*a)++; " "}"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void op(void* a_void) { " " unsigned long long* a = reinterpret_cast(a_void); " " (*a)++; " "}"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::string get_merge_sort_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " char* lhs = reinterpret_cast(lhs_void); " " char* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_UINT8: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned char* lhs = reinterpret_cast(lhs_void); " " unsigned char* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_INT16: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " short* lhs = reinterpret_cast(lhs_void); " " short* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_UINT16: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned short* lhs = reinterpret_cast(lhs_void); " " unsigned short* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " int* lhs = reinterpret_cast(lhs_void); " " int* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned int* lhs = reinterpret_cast(lhs_void); " " unsigned int* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " long long* lhs = reinterpret_cast(lhs_void); " " long long* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned long long* lhs = reinterpret_cast(lhs_void); " " unsigned long long* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " float* lhs = reinterpret_cast(lhs_void); " " float* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " double* lhs = reinterpret_cast(lhs_void); " " double* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT16: return "#include \n" "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " __half* lhs = reinterpret_cast<__half*>(lhs_void); " " __half* rhs = reinterpret_cast<__half*>(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs < *rhs; " "}"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::string get_unique_by_key_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " char* lhs = reinterpret_cast(lhs_void); " " char* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_UINT8: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned char* lhs = reinterpret_cast(lhs_void); " " unsigned char* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_INT16: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " short* lhs = reinterpret_cast(lhs_void); " " short* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_UINT16: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned short* lhs = reinterpret_cast(lhs_void); " " unsigned short* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " int* lhs = reinterpret_cast(lhs_void); " " int* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned int* lhs = reinterpret_cast(lhs_void); " " unsigned int* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " long long* lhs = reinterpret_cast(lhs_void); " " long long* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " unsigned long long* lhs = reinterpret_cast(lhs_void); " " unsigned long long* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT32: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " float* lhs = reinterpret_cast(lhs_void); " " float* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT64: return "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " double* lhs = reinterpret_cast(lhs_void); " " double* rhs = reinterpret_cast(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; case cccl_type_enum::CCCL_FLOAT16: return "#include \n" "extern \"C\" __device__ void op(void* lhs_void, void* rhs_void, void* result_void) { " " __half* lhs = reinterpret_cast<__half*>(lhs_void); " " __half* rhs = reinterpret_cast<__half*>(rhs_void); " " bool* result = reinterpret_cast(result_void); " " *result = *lhs == *rhs; " "}"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::string get_unary_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " char* a = reinterpret_cast(a_void); " " char* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " int* a = reinterpret_cast(a_void); " " int* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " unsigned int* a = reinterpret_cast(a_void); " " unsigned int* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " long long* a = reinterpret_cast(a_void); " " long long* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " unsigned long long* a = reinterpret_cast(a_void); " " unsigned long long* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_FLOAT32: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " float* a = reinterpret_cast(a_void); " " float* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_FLOAT64: return "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " double* a = reinterpret_cast(a_void); " " double* result = reinterpret_cast(result_void); " " *result = 2 * *a; " "}"; case cccl_type_enum::CCCL_FLOAT16: return "#include \n" "extern \"C\" __device__ void op(void* a_void, void* result_void) { " " __half* a = reinterpret_cast<__half*>(a_void); " " __half* result = reinterpret_cast<__half*>(result_void); " " *result = __float2half(2.0f) * (*a); " "}"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::string get_radix_sort_decomposer_op(cccl_type_enum t) { switch (t) { case cccl_type_enum::CCCL_INT8: return "extern \"C\" __device__ void* op(void* key_void) { " " char* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_UINT8: return "extern \"C\" __device__ void* op(void* key_void) { " " unsigned char* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_INT16: return "extern \"C\" __device__ void* op(void* key_void) { " " short* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_UINT16: return "extern \"C\" __device__ void* op(void* key_void) { " " unsigned short* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_INT32: return "extern \"C\" __device__ void* op(void* key_void) { " " int* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_UINT32: return "extern \"C\" __device__ void* op(void* key_void) { " " unsigned int* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_INT64: return "extern \"C\" __device__ void* op(void* key_void) { " " long long* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_UINT64: return "extern \"C\" __device__ void* op(void* key_void) { " " unsigned long long* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_FLOAT32: return "extern \"C\" __device__ void* op(void* key_void) { " " float* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_FLOAT64: return "extern \"C\" __device__ void* op(void* key_void) { " " double* key = reinterpret_cast(key_void); " " return key; " "};"; case cccl_type_enum::CCCL_FLOAT16: return "#include \n" "extern \"C\" __device__ void* op(void* key_void) { " " __half* key = reinterpret_cast<__half*>(key_void); " " return key; " "};"; default: throw std::runtime_error("Unsupported type"); } return ""; } inline std::pair get_three_way_partition_ops(cccl_type_enum t, int compare_to) { std::string less_op_src = std::format( "#include \n" "extern \"C\" __device__ void less_op(void* x_void, void* out_void) {{ " " {0}* x = reinterpret_cast<{0}*>(x_void); " " bool* out = reinterpret_cast(out_void); " " *out = *x < static_cast<{0}>({1}); " "}}", type_enum_to_name(t), compare_to); std::string greater_or_equal_op_src = std::format( "#include \n" "extern \"C\" __device__ void greater_op(void* x_void, void* out_void) {{ " " {0}* x = reinterpret_cast<{0}*>(x_void); " " bool* out = reinterpret_cast(out_void); " " *out = *x >= static_cast<{0}>({1}); " "}}", type_enum_to_name(t), compare_to); return {std::move(less_op_src), std::move(greater_or_equal_op_src)}; } template struct pointer_t { T* ptr{}; size_t size{}; pointer_t(std::size_t num_items) { REQUIRE(cudaSuccess == cudaMalloc(&ptr, num_items * sizeof(T))); size = num_items; } pointer_t(const std::vector& vec) { REQUIRE(cudaSuccess == cudaMalloc(&ptr, vec.size() * sizeof(T))); REQUIRE(cudaSuccess == cudaMemcpy(ptr, vec.data(), vec.size() * sizeof(T), cudaMemcpyHostToDevice)); size = vec.size(); } pointer_t() = default; ~pointer_t() { if (ptr) { REQUIRE(cudaSuccess == cudaFree(ptr)); ptr = nullptr; } } T operator[](int i) const { T value{}; REQUIRE(cudaSuccess == cudaMemcpy(&value, ptr + i, sizeof(T), cudaMemcpyDeviceToHost)); return value; } operator cccl_iterator_t() { cccl_iterator_t it; it.size = sizeof(T); it.alignment = alignof(T); it.type = cccl_iterator_kind_t::CCCL_POINTER; it.state = ptr; it.value_type = get_type_info(); it.advance = {}; it.dereference = {}; return it; } operator std::vector() const { std::vector vec(size); REQUIRE(cudaSuccess == cudaMemcpy(vec.data(), ptr, sizeof(T) * size, cudaMemcpyDeviceToHost)); return vec; } }; // std::vector cannot provide the contiguous storage needed by pointer_t. // Use byte storage for Boolean inputs and outputs while describing it to the C // API as its corresponding primitive type. inline cccl_iterator_t make_boolean_iterator(pointer_t& storage) { static_assert(sizeof(bool) == sizeof(uint8_t)); static_assert(alignof(bool) == alignof(uint8_t)); cccl_iterator_t iterator = storage; iterator.size = sizeof(bool); iterator.alignment = alignof(bool); iterator.value_type.size = sizeof(bool); iterator.value_type.alignment = alignof(bool); iterator.value_type.type = cccl_type_enum::CCCL_BOOLEAN; return iterator; } struct operation_t { std::string name; std::string code; cccl_op_code_type code_type = CCCL_OP_LTOIR; // Default to LTO-IR for backward compatibility operation_t() = default; operation_t(std::string_view op_name, std::string_view op_code, cccl_op_code_type op_code_type = CCCL_OP_LTOIR) : name(op_name) , code(op_code) , code_type(op_code_type) {} operator cccl_op_t() { cccl_op_t op; op.type = cccl_op_kind_t::CCCL_STATELESS; op.name = name.c_str(); op.code = code.c_str(); op.code_size = code.size(); op.code_type = code_type; op.size = 1; op.alignment = 1; op.state = nullptr; op.extra_ltoirs = nullptr; op.extra_ltoir_sizes = nullptr; op.num_extra_ltoirs = 0; return op; } }; template struct stateful_operation_t { OpT op_state; std::string name; std::string code; stateful_operation_t(const OpT& state, std::string_view op_name, std::string_view op_code) : op_state(state) , name(op_name) , code(op_code) {} operator cccl_op_t() { cccl_op_t op; op.type = cccl_op_kind_t::CCCL_STATEFUL; op.size = sizeof(OpT); op.alignment = alignof(OpT); op.state = &op_state; op.name = name.c_str(); op.code = code.c_str(); op.code_size = code.size(); op.code_type = CCCL_OP_LTOIR; // Stateful operations always use LTO-IR op.extra_ltoirs = nullptr; op.extra_ltoir_sizes = nullptr; op.num_extra_ltoirs = 0; return op; } }; inline operation_t make_operation(std::string_view name, const std::string& code) { return operation_t{name, compile(code), CCCL_OP_LTOIR}; } inline operation_t make_cpp_operation(std::string_view name, const std::string& cpp_code) { return operation_t{name, cpp_code, CCCL_OP_CPP_SOURCE}; } template stateful_operation_t make_operation(std::string_view name, const std::string& code, OpT op) { return {op, name, compile(code)}; } // Designated initializers so this header builds against both v1's and v2's // cccl_op_t — v2 adds an extra_code_types field that v1 doesn't define, and // any unlisted field zero-inits (== nullptr for the pointer fields). static cccl_op_t make_well_known_unary_operation() { return cccl_op_t{ .type = cccl_op_kind_t::CCCL_NEGATE, .name = "", .code = "", .code_size = 0, .code_type = CCCL_OP_LTOIR, .size = 1, .alignment = 1, .state = nullptr, .extra_ltoirs = nullptr, .extra_ltoir_sizes = nullptr, .num_extra_ltoirs = 0, .extra_code_types = nullptr, }; } static cccl_op_t make_well_known_binary_operation() { return cccl_op_t{ .type = cccl_op_kind_t::CCCL_PLUS, .name = "", .code = "", .code_size = 0, .code_type = CCCL_OP_LTOIR, .size = 1, .alignment = 1, .state = nullptr, .extra_ltoirs = nullptr, .extra_ltoir_sizes = nullptr, .num_extra_ltoirs = 0, .extra_code_types = nullptr, }; } static cccl_op_t make_well_known_less_binary_predicate() { return cccl_op_t{ .type = cccl_op_kind_t::CCCL_LESS, .name = "", .code = "", .code_size = 0, .code_type = CCCL_OP_LTOIR, .size = 1, .alignment = 1, .state = nullptr, .extra_ltoirs = nullptr, .extra_ltoir_sizes = nullptr, .num_extra_ltoirs = 0, .extra_code_types = nullptr, }; } static cccl_op_t make_well_known_unique_binary_predicate() { return cccl_op_t{ .type = cccl_op_kind_t::CCCL_EQUAL_TO, .name = "", .code = "", .code_size = 0, .code_type = CCCL_OP_LTOIR, .size = 1, .alignment = 1, .state = nullptr, .extra_ltoirs = nullptr, .extra_ltoir_sizes = nullptr, .num_extra_ltoirs = 0, .extra_code_types = nullptr, }; } static cccl_op_t make_well_known_greater_equal_binary_predicate() { return cccl_op_t{ .type = cccl_op_kind_t::CCCL_GREATER_EQUAL, .name = "", .code = "", .code_size = 0, .code_type = CCCL_OP_LTOIR, .size = 1, .alignment = 1, .state = nullptr, .extra_ltoirs = nullptr, .extra_ltoir_sizes = nullptr, .num_extra_ltoirs = 0, .extra_code_types = nullptr, }; } template struct iterator_t { StateT state; std::string state_name; operation_t advance; operation_t dereference; operator cccl_iterator_t() { cccl_iterator_t it; it.size = sizeof(StateT); it.alignment = alignof(StateT); it.type = cccl_iterator_kind_t::CCCL_ITERATOR; it.advance = advance; it.dereference = dereference; it.value_type = get_type_info(); it.state = &state; return it; } }; enum class iterator_kind { INPUT = 0, OUTPUT = 1, }; template struct random_access_iterator_state_t { T* data; }; template struct counting_iterator_state_t { T value; }; template struct constant_iterator_state_t { T value; }; template struct stateless_transform_it_state { using BaseIteratorStateT = BaseIteratorStateTy; BaseIteratorStateTy base_it_state; }; template struct stateful_transform_it_state { using BaseIteratorStateT = BaseIteratorStateTy; using FunctorStateT = FunctorStateTy; BaseIteratorStateTy base_it_state; FunctorStateTy functor_state; }; struct name_source_t { std::string_view name; std::string_view def_src; }; template iterator_t make_iterator(name_source_t state, operation_t advance, operation_t dereference) { iterator_t it; it.state_name = state.name; const std::string& state_src = std::string{state.def_src}; it.advance = make_operation(advance.name, state_src + advance.code); it.dereference = make_operation(dereference.name, state_src + dereference.code); return it; } inline std::tuple make_random_access_iterator_sources( iterator_kind kind, std::string_view value_type, std::string_view iterator_state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name, std::string_view transform = "") { std::string state_def_src = std::format("struct {0} {{ {1}* data; }};\n", iterator_state_name, value_type); std::string advance_fn_def_src = std::format( "extern \"C\" __device__ void {0}(void* state, const void* offset) {{\n" " auto* typed_state = static_cast<{1}*>(state);\n" " auto offset_val = *static_cast(offset);\n" " typed_state->data += offset_val;\n" "}}", advance_fn_name, iterator_state_name); std::string dereference_fn_def_src; if (kind == iterator_kind::INPUT) { dereference_fn_def_src = std::format( "extern \"C\" __device__ void {0}(const void* state, {1}* result) {{\n" " auto* typed_state = static_cast(state);\n" " *result = (*typed_state->data){3};\n" "}}", dereference_fn_name, value_type, iterator_state_name, transform); } else { dereference_fn_def_src = std::format( "extern \"C\" __device__ void {0}(void* state, const void* x) {{\n" " auto* typed_state = static_cast<{1}*>(state);\n" " auto x_val = *static_cast(x);\n" " *typed_state->data = x_val{3};\n" "}}", dereference_fn_name, iterator_state_name, value_type, transform); } return std::make_tuple(state_def_src, advance_fn_def_src, dereference_fn_def_src); } template iterator_t> make_random_access_iterator( iterator_kind kind, std::string_view value_type, std::string prefix = "", std::string transform = "") { std::string iterator_state_name = std::format("{0}state_t", prefix); std::string advance_fn_name = std::format("{0}advance", prefix); std::string dereference_fn_name = std::format("{0}dereference", prefix); const auto& [iterator_state_def_src, advance_fn_def_src, dereference_fn_def_src] = make_random_access_iterator_sources( kind, value_type, iterator_state_name, advance_fn_name, dereference_fn_name, transform); name_source_t iterator_state = {iterator_state_name, iterator_state_def_src}; operation_t advance = {advance_fn_name, advance_fn_def_src}; operation_t dereference = {dereference_fn_name, dereference_fn_def_src}; return make_iterator>(iterator_state, advance, dereference); } inline std::tuple make_counting_iterator_sources( std::string_view value_type, std::string_view iterator_state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name) { std::string iterator_state_def_src = std::format("struct {0} {{ {1} value; }};\n", iterator_state_name, value_type); std::string advance_fn_def_src = std::format( "extern \"C\" __device__ void {0}(void* state, const void* offset) {{\n" " auto* typed_state = static_cast<{1}*>(state);\n" " auto offset_val = *static_cast(offset);\n" " typed_state->value += offset_val;\n" "}}", advance_fn_name, iterator_state_name); std::string dereference_fn_def_src = std::format( "extern \"C\" __device__ void {0}(const void* state, {2}* result) {{ \n" " auto* typed_state = static_cast(state);\n" " *result = typed_state->value;\n" "}}", dereference_fn_name, iterator_state_name, value_type); return std::make_tuple(iterator_state_def_src, advance_fn_def_src, dereference_fn_def_src); } template iterator_t> make_counting_iterator(std::string_view value_type, std::string_view prefix = "") { std::string iterator_state_name = std::format("{0}state_t", prefix); std::string advance_fn_name = std::format("{0}advance", prefix); std::string dereference_fn_name = std::format("{0}dereference", prefix); const auto& [iterator_state_src, advance_fn_def_src, dereference_fn_def_src] = make_counting_iterator_sources(value_type, iterator_state_name, advance_fn_name, dereference_fn_name); name_source_t iterator_state = {iterator_state_name, iterator_state_src}; operation_t advance = {advance_fn_name, advance_fn_def_src}; operation_t dereference = {dereference_fn_name, dereference_fn_def_src}; return make_iterator>(iterator_state, advance, dereference); } inline std::tuple make_constant_iterator_sources( std::string_view value_type, std::string_view iterator_state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name) { std::string iterator_state_src = std::format("struct {0} {{ {1} value; }};\n", iterator_state_name, value_type); std::string advance_fn_src = std::format("extern \"C\" __device__ void {0}(void* state, const void* offset) {{ }}", advance_fn_name); std::string dereference_fn_src = std::format( "extern \"C\" __device__ void {0}(const void* state, {1}* result) {{ \n" " auto* typed_state = static_cast(state);\n" " *result = typed_state->value;\n" "}}", dereference_fn_name, value_type, iterator_state_name); return std::make_tuple(iterator_state_src, advance_fn_src, dereference_fn_src); } template iterator_t> make_constant_iterator(std::string_view value_type, std::string_view prefix = "") { std::string iterator_state_name = std::format("{0}struct_t", prefix); std::string advance_fn_name = std::format("{0}advance", prefix); std::string dereference_fn_name = std::format("{0}dereference", prefix); const auto& [iterator_state_src, advance_fn_src, dereference_fn_src] = make_constant_iterator_sources(value_type, iterator_state_name, advance_fn_name, dereference_fn_name); name_source_t iterator_state = {iterator_state_name, iterator_state_src}; operation_t advance = {advance_fn_name, advance_fn_src}; operation_t dereference = {dereference_fn_name, dereference_fn_src}; return make_iterator>(iterator_state, advance, dereference); } inline std::tuple make_reverse_iterator_sources( iterator_kind kind, std::string_view value_type, std::string_view iterator_state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name, std::string_view transform = "") { std::string iterator_state_src = std::format("struct {0} {{ {1}* data; }};\n", iterator_state_name, value_type); std::string advance_fn_src = std::format( "extern \"C\" __device__ void {0}(void* state, const void* offset) {{\n" " auto* typed_state = static_cast<{1}*>(state);\n" " auto offset_val = *static_cast(offset);\n" " typed_state->data -= offset_val;\n" "}}", advance_fn_name, iterator_state_name); std::string dereference_fn_src; if (kind == iterator_kind::INPUT) { dereference_fn_src = std::format( "extern \"C\" __device__ void {0}(const void* state, {2}* result) {{\n" " auto* typed_state = static_cast(state);\n" " *result = (*typed_state->data){3};\n" "}}", dereference_fn_name, iterator_state_name, value_type, transform); } else { dereference_fn_src = std::format( "extern \"C\" __device__ void {0}(void* state, const void* x) {{\n" " auto* typed_state = static_cast<{1}*>(state);\n" " auto x_val = *static_cast(x);\n" " *typed_state->data = x_val{3};\n" "}}", dereference_fn_name, iterator_state_name, value_type, transform); } return std::make_tuple(iterator_state_src, advance_fn_src, dereference_fn_src); } inline std::tuple make_step_counting_iterator_sources( std::string_view index_ty_name, std::string_view state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name) { static constexpr std::string_view it_state_src_tmpl = R"XXX( struct {0} {{ {1} linear_id; {1} segment_size; }}; )XXX"; const std::string it_state_def_src = std::format(it_state_src_tmpl, state_name, index_ty_name); static constexpr std::string_view it_def_src_tmpl = R"XXX( extern "C" __device__ void {0}(void* state, const void* offset) {{ auto* typed_state = static_cast<{1}*>(state); auto offset_val = *static_cast(offset); typed_state->linear_id += offset_val; }} )XXX"; const std::string it_advance_fn_def_src = std::format(it_def_src_tmpl, /*0*/ advance_fn_name, state_name, index_ty_name); static constexpr std::string_view it_deref_src_tmpl = R"XXX( extern "C" __device__ void {0}(const void* state, {1}* result) {{ auto* typed_state = static_cast(state); *result = (typed_state->linear_id) * (typed_state->segment_size); }} )XXX"; const std::string it_deref_fn_def_src = std::format(it_deref_src_tmpl, dereference_fn_name, index_ty_name, state_name); return std::make_tuple(it_state_def_src, it_advance_fn_def_src, it_deref_fn_def_src); } // Host-side advance function for iterator states that have a `linear_id` member template inline void host_advance_linear_id(void* state, cccl_increment_t offset) { auto* st = reinterpret_cast(state); using Index = decltype(st->linear_id); if constexpr (std::is_signed_v) { st->linear_id += offset.signed_offset; } else { st->linear_id += offset.unsigned_offset; } } // Host-side advance for iterator states that contain a nested `base_it_state.value` template inline void host_advance_base_value(void* state, cccl_increment_t offset) { auto st = reinterpret_cast(state); using IndexT = decltype(st->base_it_state.value); if constexpr (std::is_signed_v) { st->base_it_state.value += offset.signed_offset; } else { st->base_it_state.value += offset.unsigned_offset; } } template iterator_t> make_reverse_iterator( iterator_kind kind, std::string_view value_type, std::string_view prefix = "", std::string_view transform = "") { std::string iterator_state_name = std::format("{0}struct_t", prefix); std::string advance_fn_name = std::format("{0}advance", prefix); std::string dereference_fn_name = std::format("{0}dereference", prefix); const auto& [iterator_state_src, advance_fn_src, dereference_fn_src] = make_reverse_iterator_sources( kind, value_type, iterator_state_name, advance_fn_name, dereference_fn_name, transform); name_source_t iterator_state = {iterator_state_name, iterator_state_src}; operation_t advance = {advance_fn_name, advance_fn_src}; operation_t dereference = {dereference_fn_name, dereference_fn_src}; return make_iterator>(iterator_state, advance, dereference); } inline std::tuple make_stateful_transform_input_iterator_sources( std::string_view transform_it_state_name, std::string_view transform_it_advance_fn_name, std::string_view transform_it_dereference_fn_name, std::string_view transformed_value_type, std::string_view base_value_type, name_source_t base_it_state, name_source_t base_it_advance_fn, name_source_t base_it_dereference_fn, name_source_t transform_state, name_source_t transform_op) { static constexpr std::string_view transform_it_state_src_tmpl = R"XXX( /* Define state of stateful transform operation */ {3} /* Define state of base iterator over whose values transformation is applied */ {4} struct {0} {{ {1} base_it_state; {2} functor_state; }}; )XXX"; const std::string transform_it_state_src = std::format( transform_it_state_src_tmpl, /* 0 */ transform_it_state_name, /* 1 */ base_it_state.name, /* 2 */ transform_state.name, /* 3 */ transform_state.def_src, /* 4 */ base_it_state.def_src); static constexpr std::string_view transform_it_advance_fn_src_tmpl = R"XXX( {3} extern "C" __device__ void {0}(void* transform_it_state, const void* offset) {{ auto* typed_state = static_cast<{1}*>(transform_it_state); {2}(&(typed_state->base_it_state), offset); }} )XXX"; const std::string transform_it_advance_fn_src = std::format( transform_it_advance_fn_src_tmpl, /* 0 */ transform_it_advance_fn_name, /* 1 */ transform_it_state_name, /* 2 */ base_it_advance_fn.name, /* 3 */ base_it_advance_fn.def_src); static constexpr std::string_view transform_it_dereference_fn_src_tmpl = R"XXX( {5} {6} extern "C" __device__ void {0}(const void* transform_it_state, {2}* result) {{ auto* typed_state = static_cast(transform_it_state); {7} base_result; {4}(&(typed_state->base_it_state), &base_result); *result = {3}( const_castfunctor_state)*>(&(typed_state->functor_state)), base_result ); }} )XXX"; const std::string transform_it_dereference_fn_src = std::format( transform_it_dereference_fn_src_tmpl, /* 0 */ transform_it_dereference_fn_name /* name of transform's deref function */, /* 1 */ transform_it_state_name /* name of transform's state*/, /* 2 */ transformed_value_type /* function return type name */, /* 3 */ transform_op.name /* transformation functor function name */, /* 4 */ base_it_dereference_fn.name /* deref function of base iterator */, /* 5 */ base_it_dereference_fn.def_src, /* 6 */ transform_op.def_src, /* 7 */ base_value_type); return std::make_tuple(transform_it_state_src, transform_it_advance_fn_src, transform_it_dereference_fn_src); } template auto make_stateful_transform_input_iterator( std::string_view transformed_value_type, std::string_view base_value_type, name_source_t base_it_state, name_source_t base_it_advance_fn, name_source_t base_it_dereference_fn, name_source_t transform_state, name_source_t transform_op) { static constexpr std::string_view transform_it_state_name = "stateful_transform_iterator_state_t"; static constexpr std::string_view transform_it_advance_fn_name = "advance_stateful_transform_it"; static constexpr std::string_view transform_it_dereference_fn_name = "dereference_stateful_transform_it"; const auto& [transform_it_state_src, transform_it_advance_fn_src, transform_it_dereference_fn_src] = make_stateful_transform_input_iterator_sources( transform_it_state_name, transform_it_advance_fn_name, transform_it_dereference_fn_name, transformed_value_type, base_value_type, base_it_state, base_it_advance_fn, base_it_dereference_fn, transform_state, transform_op); using HostTransformStateT = stateful_transform_it_state; auto transform_it = make_iterator( {transform_it_state_name, transform_it_state_src}, {transform_it_advance_fn_name, transform_it_advance_fn_src}, {transform_it_dereference_fn_name, transform_it_dereference_fn_src}); return transform_it; } /*! @brief Generate source code with definitions for state of transformed iterator and functions to operator on it */ inline std::tuple make_stateless_transform_input_iterator_sources( std::string_view transform_it_state_name, std::string_view transform_it_advance_fn_name, std::string_view transform_it_dereference_fn_name, std::string_view transformed_value_type, std::string_view base_value_type, name_source_t base_it_state, name_source_t base_it_advance_fn, name_source_t base_it_dereference_fn, name_source_t transform_op) { static constexpr std::string_view transform_it_state_src_tmpl = R"XXX( /* Define state of base iterator over whose values transformation is applied */ {2} struct {0} {{ {1} base_it_state; }}; )XXX"; const std::string transform_it_state_src = std::format( transform_it_state_src_tmpl, /* 0 */ transform_it_state_name, /* 1 */ base_it_state.name, /* 2 */ base_it_state.def_src); static constexpr std::string_view transform_it_advance_fn_src_tmpl = R"XXX( {3} extern "C" __device__ void {0}(void *transform_it_state, const void* offset) {{ auto* typed_state = static_cast<{1}*>(transform_it_state); {2}(&(typed_state->base_it_state), offset); }} )XXX"; const std::string transform_it_advance_fn_src = std::format( transform_it_advance_fn_src_tmpl, /* 0 */ transform_it_advance_fn_name, /* 1 */ transform_it_state_name, /* 2 */ base_it_advance_fn.name, /* 3 */ base_it_advance_fn.def_src); static constexpr std::string_view transform_it_dereference_fn_src_tmpl = R"XXX( {5} {6} extern "C" __device__ void {0}({1} *transform_it_state, {2}* result) {{ {7} base_result; {4}(&(transform_it_state->base_it_state), &base_result); *result = {3}(base_result); }} )XXX"; const std::string transform_it_dereference_fn_src = std::format( transform_it_dereference_fn_src_tmpl, /* 0 */ transform_it_dereference_fn_name /* name of transform's deref function */, /* 1 */ transform_it_state_name /* name of transform's state*/, /* 2 */ transformed_value_type /* function return type name */, /* 3 */ transform_op.name /* transformation functor function name */, /* 4 */ base_it_dereference_fn.name /* deref function of base iterator */, /* 5 */ base_it_dereference_fn.def_src, /* 6 */ transform_op.def_src, /* 7 */ base_value_type); return std::make_tuple(transform_it_state_src, transform_it_advance_fn_src, transform_it_dereference_fn_src); } template auto make_stateless_transform_input_iterator( std::string_view transformed_value_type, std::string_view base_value_type, name_source_t base_it_state, name_source_t base_it_advance_fn, name_source_t base_it_dereference_fn, name_source_t transform_op) { static constexpr std::string_view transform_it_state_name = "stateless_transform_iterator_state_t"; static constexpr std::string_view transform_it_advance_fn_name = "advance_stateless_transform_it"; static constexpr std::string_view transform_it_deref_fn_name = "dereference_stateless_transform_it"; const auto& [transform_it_state_src, transform_it_advance_fn_src, transform_it_deref_fn_src] = make_stateless_transform_input_iterator_sources( transform_it_state_name, transform_it_advance_fn_name, transform_it_deref_fn_name, transformed_value_type, base_value_type, base_it_state, base_it_advance_fn, base_it_dereference_fn, transform_op); using HostTransformStateT = stateless_transform_it_state; auto transform_it = make_iterator( {transform_it_state_name, transform_it_state_src}, {transform_it_advance_fn_name, transform_it_advance_fn_src}, {transform_it_deref_fn_name, transform_it_deref_fn_src}); return transform_it; } inline std::tuple make_discard_iterator_sources( iterator_kind kind, std::string_view value_type, std::string_view iterator_state_name, std::string_view advance_fn_name, std::string_view dereference_fn_name) { std::string state_def_src = std::format("struct {0} {{ {1}* data; }};\n", iterator_state_name, value_type); std::string advance_fn_def_src = std::format( "extern \"C\" __device__ void {0}(void* /*state*/, const void* /*offset*/) {{\n" "}}", advance_fn_name, iterator_state_name); std::string dereference_fn_def_src; if (kind == iterator_kind::INPUT) { dereference_fn_def_src = std::format( "extern \"C\" __device__ void {0}(const void* /*state*/, {2}* /*result*/) {{\n" "}}", dereference_fn_name, iterator_state_name, value_type); } else { dereference_fn_def_src = std::format( "extern \"C\" __device__ void {0}(void* /*state*/, const void* /*x*/) {{\n" "}}", dereference_fn_name, iterator_state_name, value_type); } return std::make_tuple(state_def_src, advance_fn_def_src, dereference_fn_def_src); } template auto make_discard_iterator(iterator_kind kind, std::string_view value_type, std::string prefix = "") { std::string iterator_state_name = std::format("{0}struct_t", prefix); std::string advance_fn_name = std::format("{0}advance", prefix); std::string dereference_fn_name = std::format("{0}dereference", prefix); const auto& [iterator_state_src, advance_fn_src, dereference_fn_src] = make_discard_iterator_sources(kind, value_type, iterator_state_name, advance_fn_name, dereference_fn_name); name_source_t iterator_state = {iterator_state_name, iterator_state_src}; operation_t advance = {advance_fn_name, advance_fn_src}; operation_t dereference = {dereference_fn_name, dereference_fn_src}; return make_iterator>(iterator_state, advance, dereference); } template struct value_t { T value; value_t(T value) : value(value) {} operator cccl_value_t() { cccl_value_t v; v.type = get_type_info(); v.state = &value; return v; } };