// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception #pragma once #include #if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC) # pragma GCC system_header #elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG) # pragma clang system_header #elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC) # pragma system_header #endif // no system header #include #include #include #include #include CUB_NAMESPACE_BEGIN namespace detail::params { // ===================================================================== // get_param — unified segment parameter access // ===================================================================== //! @brief Returns the value of an argument for a given segment index. //! //! @param[in] __arg Argument or argument wrapper to read. //! @param[in] __index Segment index to read for sequence arguments. //! @return The single argument value, or the sequence element at the given index. _CCCL_TEMPLATE(class _Tp, class _SegmentIndexT) _CCCL_REQUIRES((!::cuda::args::__is_wrapper_v<::cuda::std::remove_cvref_t<_Tp>>) ) [[nodiscard]] _CCCL_HOST_DEVICE constexpr auto get_param(_Tp&& __arg, [[maybe_unused]] _SegmentIndexT __index) noexcept { if constexpr (::cuda::args::__traits<::cuda::std::remove_cvref_t<_Tp>>::is_single_value) { return __arg; } else { return __arg[__index]; } } template [[nodiscard]] _CCCL_HOST_DEVICE constexpr auto get_param(const ::cuda::args::constant<_Value, _Tp>& __arg, [[maybe_unused]] _SegmentIndexT __index) noexcept { return ::cuda::args::__unwrap(__arg); } template [[nodiscard]] _CCCL_HOST_DEVICE constexpr auto get_param(const ::cuda::args::immediate<_Arg, _StaticBounds>& __arg, [[maybe_unused]] _SegmentIndexT __index) noexcept { return ::cuda::args::__unwrap(__arg); } template [[nodiscard]] _CCCL_HOST_DEVICE constexpr auto get_param(const ::cuda::args::deferred<_Arg, _StaticBounds>& __arg, [[maybe_unused]] _SegmentIndexT __index) noexcept { return ::cuda::args::__unwrap(__arg); } template [[nodiscard]] _CCCL_HOST_DEVICE constexpr auto get_param(const ::cuda::args::deferred_sequence<_Arg, _StaticBounds>& __arg, _SegmentIndexT __index) noexcept { return ::cuda::args::__unwrap(__arg)[__index]; } // ===================================================================== // Discrete parameter support // ===================================================================== //! @brief Specifies a list of supported options for a parameter. template struct supported_options { static constexpr ::cuda::std::size_t count = sizeof...(Options); }; //! @brief Static discrete parameter — a single compile-time value that is also its only supported option. //! //! Holds no runtime value, so it cannot be put into a state that disagrees with its supported option, and //! @c dispatch_impl therefore always matches it. This is the safe representation for a compile-time-fixed discrete //! parameter (e.g. a statically known top-k selection direction): modeling such a parameter with a runtime value //! instead would risk that value silently disagreeing with the supported option (a no-op dispatch unless //! @c CCCL_ENABLE_ASSERTIONS is set). template struct static_discrete_param { using value_type = T; using supported_options_t = supported_options; template [[nodiscard]] _CCCL_HOST_DEVICE constexpr T get_param(SegmentIndexT) const noexcept { return Value; } }; // ===================================================================== // Discrete dispatch // ===================================================================== //! @brief Translates a runtime parameter value into a compile-time constant by matching //! against a list of supported options. //! //! @param[in] val Runtime value to match. //! @param[in] __supported_options Supported values for the parameter. //! @param[in] f Functor invoked with the matched compile-time constant. //! @return `true` if the value matches one of the supported options. template [[nodiscard]] _CCCL_HOST_DEVICE bool dispatch_impl(T val, [[maybe_unused]] supported_options __supported_options, Functor&& f) { const bool match_found = ((val == Opts ? (f(::cuda::std::integral_constant{}), true) : false) || ...); _CCCL_ASSERT(match_found, "The given runtime parameter value is not in the supported list"); return match_found; } //! @brief Dispatcher that resolves a discrete parameter to a compile-time constant //! and invokes a functor with the matched option. //! //! @param[in] param Discrete parameter to resolve. //! @param[in] segment_id Segment index to read from `param`. //! @param[in] f Functor invoked with the matched compile-time constant. //! @return `true` if the parameter value matches one of its supported options. template [[nodiscard]] _CCCL_HOST_DEVICE bool dispatch_discrete(ParamT param, SegmentIndexT segment_id, Functor&& f) { using supported_list = typename ParamT::supported_options_t; auto param_value = param.get_param(segment_id); return CUB_NS_QUALIFIER::detail::params::dispatch_impl( param_value, supported_list{}, ::cuda::std::forward(f)); } } // namespace detail::params CUB_NAMESPACE_END