Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
6eeb08f
exposes batched top-k for small segs
elstehle Jun 9, 2026
8c8e07a
Document DeviceBatchedTopK determinism, tie-break, and output ordering
elstehle Jun 14, 2026
536e0e2
extends top-k requirements page for DeviceBatchedTopK
elstehle Jun 15, 2026
bad1e4c
improves docs and examples
elstehle Jun 16, 2026
ca103b6
better cuda::args docs
elstehle Jun 17, 2026
fbd66d7
drops standalone examples for follow-up pr
elstehle Jun 17, 2026
b6c8113
addresses review comments
elstehle Jun 17, 2026
8111d1b
refines docs
elstehle Jun 22, 2026
2f7f695
introduces the both-or-nothing rule for determinism and tie-break
elstehle Jun 23, 2026
d851e22
adds documentation on function parameters
elstehle Jun 23, 2026
c777859
validates unwrapped parameters
elstehle Jun 23, 2026
6cb0864
clarifies support for plain integrals
elstehle Jun 23, 2026
b9ffd43
fixes copydetails
elstehle Jun 23, 2026
ef9d2b6
Merge remote-tracking branch 'upstream/main' into topk/expose-device-…
elstehle Jun 23, 2026
3625924
Merge remote-tracking branch 'upstream/main' into topk/expose-device-…
elstehle Jun 24, 2026
c1226fb
Merge remote-tracking branch 'upstream/main' into topk/expose-device-…
elstehle Jun 27, 2026
4cbb929
adds api test for unwrapped integrals
elstehle Jun 27, 2026
6b0e679
addresses review comments
elstehle Jun 27, 2026
800ce18
Merge remote-tracking branch 'upstream/main' into topk/expose-device-…
elstehle Jun 30, 2026
840ebdd
adds to the umbrella cub.cuh and device-wide docs
elstehle Jun 30, 2026
484c418
const-rev env and docs on in/out overlapping
elstehle Jun 30, 2026
4c7a6e1
improves api examples; adds mis-aligned temp storage test
elstehle Jun 30, 2026
57000c6
Merge remote-tracking branch 'upstream/main' into topk/expose-device-…
elstehle Jun 30, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cub/cub/cub.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@

// Device
#include <cub/device/device_adjacent_difference.cuh>
#include <cub/device/device_batched_topk.cuh>
#include <cub/device/device_copy.cuh>
#include <cub/device/device_find.cuh>
#include <cub/device/device_for.cuh>
Expand Down
1,132 changes: 1,132 additions & 0 deletions cub/cub/device/device_batched_topk.cuh

Large diffs are not rendered by default.

11 changes: 8 additions & 3 deletions cub/cub/device/device_topk.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -69,11 +69,11 @@ CUB_RUNTIME_FUNCTION static cudaError_t dispatch_topk(
using requested_determinism_t =
::cuda::std::execution::__query_result_or_t<requirements_t,
::cuda::execution::determinism::__get_determinism_t,
::cuda::execution::determinism::run_to_run_t>;
::cuda::execution::determinism::gpu_to_gpu_t>;
using requested_order_t =
::cuda::std::execution::__query_result_or_t<requirements_t,
::cuda::execution::output_ordering::__get_output_ordering_t,
::cuda::execution::output_ordering::sorted_t>;
::cuda::execution::output_ordering::stable_sorted_t>;
constexpr auto is_determinism_not_guaranteed =
::cuda::std::is_same_v<requested_determinism_t, ::cuda::execution::determinism::not_guaranteed_t>;
constexpr auto is_output_order_unsorted =
Expand All @@ -84,6 +84,11 @@ CUB_RUNTIME_FUNCTION static cudaError_t dispatch_topk(
"cub::DeviceTopK only supports the case where determinism is not guaranteed and output order is "
"unsorted.");

// TODO (elstehle): align requirement validation with cub::DeviceBatchedTopK in CCCL 4.0. cub::DeviceTopK does not
// yet inspect cuda::execution::tie_break, so it still accepts requirement combinations that cub::DeviceBatchedTopK
// rejects. It should enforce that determinism and tie_break are requested together (or both omitted to take the
// default) and that an explicit tie_break requires cuda::execution::determinism::gpu_to_gpu.

// Query relevant properties from the environment
auto stream = ::cuda::__call_or(::cuda::get_stream, ::cuda::stream_ref{cudaStream_t{}}, env);

Expand Down Expand Up @@ -167,7 +172,7 @@ CUB_RUNTIME_FUNCTION static cudaError_t dispatch_topk_hub(
//! The result of ``DeviceTopK`` is governed by two orthogonal execution requirements: *which* items are selected
//! (``cuda::execution::determinism``, optionally refined by ``cuda::execution::tie_break``) and the order in which
//! they are written (``cuda::execution::output_ordering``). When the caller does not opt out, the committed default
//! is the most reproducible behavior: deterministic results (``cuda::execution::determinism::run_to_run``), ties
//! is the most reproducible behavior: deterministic results (``cuda::execution::determinism::gpu_to_gpu``), ties
//! resolved toward the smaller (lower) source index (``cuda::execution::tie_break::prefer_smaller_index``), and
//! stable-sorted output (``cuda::execution::output_ordering::stable_sorted``). Callers opt *out* of these guarantees
//! to obtain faster implementations.
Expand Down
6 changes: 5 additions & 1 deletion cub/cub/device/dispatch/kernels/kernel_batched_topk.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,11 @@ private:

public:
// TODO (elstehle): extend support for variable-size segments
static_assert(selected_index >= 0, "No valid policy found for one-worker-per-segment approach");
static_assert(selected_index >= 0,
"cub::DeviceBatchedTopK currently supports only segments small enough to be processed by a single "
"thread block (one worker per segment). No policy could cover the statically-known maximum segment "
"size within the shared-memory limit. Reduce the maximum segment size encoded in the segment-size "
"argument annotation.");
static constexpr policy_t policy = {
active_policy.worker_per_segment_policies[selected_index], active_policy.multi_worker_per_segment_policy};

Expand Down
317 changes: 317 additions & 0 deletions cub/test/catch2_test_device_batched_topk_api.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,317 @@
// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception

#include "insert_nested_NVTX_range_guard.h"

#include <cub/device/device_batched_topk.cuh>

#include <thrust/detail/raw_pointer_cast.h>
#include <thrust/device_vector.h>
#include <thrust/host_vector.h>
#include <thrust/sort.h>

#include <cuda/__execution/determinism.h>
#include <cuda/__execution/output_ordering.h>
#include <cuda/__execution/require.h>
#include <cuda/argument>
#include <cuda/iterator>
#include <cuda/std/__execution/env.h>
#include <cuda/std/functional>

#include <c2h/catch2_test_helper.h>

C2H_TEST("cub::DeviceBatchedTopK::MaxKeys temp-storage API example", "[batched_topk][device]")
{
// example-begin batched-topk-max-keys
constexpr int num_segments = 2;
constexpr int segment_size = 8;
constexpr int k = 3;

auto keys_in = thrust::device_vector<int>{5, -3, 1, 7, 8, 2, 4, 6, /**/ 0, 9, 3, 2, 1, 8, 7, 4};
auto keys_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);
Comment thread
elstehle marked this conversation as resolved.

// Per-segment iterators: d_keys_in[s] yields an iterator to the start of segment s.
auto d_keys_in =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_in.data())), segment_size);
auto d_keys_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_out.data())), k);

// Argument annotations: a small, compile-time segment size and k, plus the runtime segment count and item-count
// bound.
constexpr auto segment_sizes = cuda::args::constant<segment_size>{};
constexpr auto k_arg = cuda::args::constant<k>{};
auto num_segs = cuda::args::immediate{cuda::std::int64_t{num_segments}};

// Top-k output is unordered and may be non-deterministic; this must be acknowledged via the environment.
auto env = cuda::std::execution::env{cuda::execution::require(
Comment thread
elstehle marked this conversation as resolved.
cuda::execution::determinism::not_guaranteed,
cuda::execution::tie_break::unspecified,
cuda::execution::output_ordering::unsorted)};

// Query temporary storage requirements
size_t temp_storage_bytes = 0;
auto error = cub::DeviceBatchedTopK::MaxKeys(
nullptr, temp_storage_bytes, d_keys_in, d_keys_out, segment_sizes, k_arg, num_segs, env);

// Allocate temporary storage and run
thrust::device_vector<char> temp_storage(temp_storage_bytes, thrust::no_init);
error = cub::DeviceBatchedTopK::MaxKeys(
thrust::raw_pointer_cast(temp_storage.data()),
Comment thread
elstehle marked this conversation as resolved.
temp_storage_bytes,
d_keys_in,
d_keys_out,
segment_sizes,
k_arg,
num_segs,
env);
// Each segment's k largest keys are written to keys_out in unspecified order. The result set is fixed,
// shown here sorted per segment:
auto expected_result_set = thrust::device_vector<int>{8, 7, 6, /* segment 0 */ 9, 8, 7 /* segment 1 */};
// example-end batched-topk-max-keys

REQUIRE(error == cudaSuccess);
Comment thread
elstehle marked this conversation as resolved.
// keys_out is unordered, so sort each segment (descending) before comparing against the expected set.
thrust::sort(keys_out.begin(), keys_out.begin() + k, cuda::std::greater<int>{});
thrust::sort(keys_out.begin() + k, keys_out.begin() + 2 * k, cuda::std::greater<int>{});
REQUIRE(keys_out == expected_result_set);
}

C2H_TEST("cub::DeviceBatchedTopK::MinKeys temp-storage API example", "[batched_topk][device]")
{
// example-begin batched-topk-min-keys
constexpr int num_segments = 2;
constexpr int segment_size = 8;
constexpr int k = 3;

auto keys_in = thrust::device_vector<int>{5, -3, 1, 7, 8, 2, 4, 6, /**/ 0, 9, 3, 2, 1, 8, 7, 4};
auto keys_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);

auto d_keys_in =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_in.data())), segment_size);
auto d_keys_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_out.data())), k);

constexpr auto segment_sizes = cuda::args::constant<segment_size>{};
constexpr auto k_arg = cuda::args::constant<k>{};
auto num_segs = cuda::args::immediate{cuda::std::int64_t{num_segments}};
auto env = cuda::std::execution::env{cuda::execution::require(
cuda::execution::determinism::not_guaranteed,
cuda::execution::tie_break::unspecified,
cuda::execution::output_ordering::unsorted)};

size_t temp_storage_bytes = 0;
auto error = cub::DeviceBatchedTopK::MinKeys(
nullptr, temp_storage_bytes, d_keys_in, d_keys_out, segment_sizes, k_arg, num_segs, env);
thrust::device_vector<char> temp_storage(temp_storage_bytes, thrust::no_init);
error = cub::DeviceBatchedTopK::MinKeys(
thrust::raw_pointer_cast(temp_storage.data()),
temp_storage_bytes,
d_keys_in,
d_keys_out,
segment_sizes,
k_arg,
num_segs,
env);
// Each segment's k smallest keys are written to keys_out in unspecified order. The result set is fixed,
// shown here sorted per segment:
auto expected_result_set = thrust::device_vector<int>{-3, 1, 2, /* segment 0 */ 0, 1, 2 /* segment 1 */};
// example-end batched-topk-min-keys

REQUIRE(error == cudaSuccess);
// keys_out is unordered, so sort each segment (ascending) before comparing against the expected set.
thrust::sort(keys_out.begin(), keys_out.begin() + k);
thrust::sort(keys_out.begin() + k, keys_out.begin() + 2 * k);
REQUIRE(keys_out == expected_result_set);
}

C2H_TEST("cub::DeviceBatchedTopK::MaxPairs temp-storage API example", "[batched_topk][device]")
{
// example-begin batched-topk-max-pairs
constexpr int num_segments = 2;
constexpr int segment_size = 8;
constexpr int k = 3;

auto keys_in = thrust::device_vector<int>{5, -3, 1, 7, 8, 2, 4, 6, /**/ 0, 9, 3, 2, 1, 8, 7, 4};
auto keys_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);
auto values_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);

auto d_keys_in =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_in.data())), segment_size);
auto d_keys_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_out.data())), k);
// Input values are the per-segment item indices [0, segment_size).
auto d_values_in = cuda::make_constant_iterator(cuda::make_counting_iterator(0));
auto d_values_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(values_out.data())), k);

constexpr auto segment_sizes = cuda::args::constant<segment_size>{};
constexpr auto k_arg = cuda::args::constant<k>{};
auto num_segs = cuda::args::immediate{cuda::std::int64_t{num_segments}};
auto env = cuda::std::execution::env{cuda::execution::require(
cuda::execution::determinism::not_guaranteed,
cuda::execution::tie_break::unspecified,
cuda::execution::output_ordering::unsorted)};

size_t temp_storage_bytes = 0;
auto error = cub::DeviceBatchedTopK::MaxPairs(
nullptr, temp_storage_bytes, d_keys_in, d_keys_out, d_values_in, d_values_out, segment_sizes, k_arg, num_segs, env);
thrust::device_vector<char> temp_storage(temp_storage_bytes, thrust::no_init);
error = cub::DeviceBatchedTopK::MaxPairs(
thrust::raw_pointer_cast(temp_storage.data()),
temp_storage_bytes,
d_keys_in,
d_keys_out,
d_values_in,
d_values_out,
segment_sizes,
k_arg,
num_segs,
env);
// keys_out holds each segment's k largest keys. The key set is fixed (shown here sorted per segment). For
// keys that tie, which equal element's value is returned is unspecified.
auto expected_result_set = thrust::device_vector<int>{8, 7, 6, /* segment 0 */ 9, 8, 7 /* segment 1 */};
// example-end batched-topk-max-pairs

REQUIRE(error == cudaSuccess);

// Each returned value is the source index of its key within the segment. Check that every value indexes
// back to the input element whose key was selected.
thrust::host_vector<int> h_keys_in(keys_in);
thrust::host_vector<int> h_keys_out(keys_out);
thrust::host_vector<int> h_values_out(values_out);
for (int s = 0; s < num_segments; ++s)
{
for (int j = 0; j < k; ++j)
{
const int idx = s * k + j;
const int v = h_values_out[idx];
REQUIRE(v >= 0);
REQUIRE(v < segment_size);
REQUIRE(h_keys_in[s * segment_size + v] == h_keys_out[idx]);
}
}

// keys_out is unordered, so sort each segment (descending) before comparing against the expected set.
thrust::sort(keys_out.begin(), keys_out.begin() + k, cuda::std::greater<int>{});
thrust::sort(keys_out.begin() + k, keys_out.begin() + 2 * k, cuda::std::greater<int>{});
REQUIRE(keys_out == expected_result_set);
}

C2H_TEST("cub::DeviceBatchedTopK::MinPairs temp-storage API example", "[batched_topk][device]")
{
// example-begin batched-topk-min-pairs
constexpr int num_segments = 2;
constexpr int segment_size = 8;
constexpr int k = 3;

auto keys_in = thrust::device_vector<int>{5, -3, 1, 7, 8, 2, 4, 6, /**/ 0, 9, 3, 2, 1, 8, 7, 4};
auto keys_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);
auto values_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);

auto d_keys_in =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_in.data())), segment_size);
auto d_keys_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_out.data())), k);
auto d_values_in = cuda::make_constant_iterator(cuda::make_counting_iterator(0));
auto d_values_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(values_out.data())), k);

constexpr auto segment_sizes = cuda::args::constant<segment_size>{};
constexpr auto k_arg = cuda::args::constant<k>{};
auto num_segs = cuda::args::immediate{cuda::std::int64_t{num_segments}};
auto env = cuda::std::execution::env{cuda::execution::require(
cuda::execution::determinism::not_guaranteed,
cuda::execution::tie_break::unspecified,
cuda::execution::output_ordering::unsorted)};

size_t temp_storage_bytes = 0;
auto error = cub::DeviceBatchedTopK::MinPairs(
nullptr, temp_storage_bytes, d_keys_in, d_keys_out, d_values_in, d_values_out, segment_sizes, k_arg, num_segs, env);
thrust::device_vector<char> temp_storage(temp_storage_bytes, thrust::no_init);
error = cub::DeviceBatchedTopK::MinPairs(
thrust::raw_pointer_cast(temp_storage.data()),
temp_storage_bytes,
d_keys_in,
d_keys_out,
d_values_in,
d_values_out,
segment_sizes,
k_arg,
num_segs,
env);
// keys_out holds each segment's k smallest keys. The key set is fixed (shown here sorted per segment). For
// keys that tie, which equal element's value is returned is unspecified.
auto expected_result_set = thrust::device_vector<int>{-3, 1, 2, /* segment 0 */ 0, 1, 2 /* segment 1 */};
// example-end batched-topk-min-pairs

REQUIRE(error == cudaSuccess);

// Each returned value is the source index of its key within the segment. Check that every value indexes
// back to the input element whose key was selected.
thrust::host_vector<int> h_keys_in(keys_in);
thrust::host_vector<int> h_keys_out(keys_out);
thrust::host_vector<int> h_values_out(values_out);
for (int s = 0; s < num_segments; ++s)
{
for (int j = 0; j < k; ++j)
{
const int idx = s * k + j;
const int v = h_values_out[idx];
REQUIRE(v >= 0);
REQUIRE(v < segment_size);
REQUIRE(h_keys_in[s * segment_size + v] == h_keys_out[idx]);
}
}

// keys_out is unordered, so sort each segment (ascending) before comparing against the expected set.
thrust::sort(keys_out.begin(), keys_out.begin() + k);
thrust::sort(keys_out.begin() + k, keys_out.begin() + 2 * k);
REQUIRE(keys_out == expected_result_set);
}

// The temporary storage size requirement must not assume a particular base-pointer alignment (the public contract
// states that no special alignment is required). Over-allocate by one byte and offset the base pointer.
C2H_TEST("cub::DeviceBatchedTopK::MaxKeys handles a misaligned temporary storage pointer", "[batched_topk][device]")
{
constexpr int num_segments = 2;
constexpr int segment_size = 8;
constexpr int k = 3;

auto keys_in = thrust::device_vector<int>{5, -3, 1, 7, 8, 2, 4, 6, /**/ 0, 9, 3, 2, 1, 8, 7, 4};
auto keys_out = thrust::device_vector<int>(num_segments * k, thrust::no_init);

auto d_keys_in =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_in.data())), segment_size);
auto d_keys_out =
cuda::make_strided_iterator(cuda::make_counting_iterator(thrust::raw_pointer_cast(keys_out.data())), k);

constexpr auto segment_sizes = cuda::args::constant<segment_size>{};
constexpr auto k_arg = cuda::args::constant<k>{};
auto num_segs = cuda::args::immediate{cuda::std::int64_t{num_segments}};
auto env = cuda::std::execution::env{cuda::execution::require(
cuda::execution::determinism::not_guaranteed,
cuda::execution::tie_break::unspecified,
cuda::execution::output_ordering::unsorted)};

size_t temp_storage_bytes = 0;
auto error = cub::DeviceBatchedTopK::MaxKeys(
nullptr, temp_storage_bytes, d_keys_in, d_keys_out, segment_sizes, k_arg, num_segs, env);
REQUIRE(error == cudaSuccess);

// Allocate one extra byte and offset the base pointer by one to misalign it.
thrust::device_vector<char> temp_storage(temp_storage_bytes + 1, thrust::no_init);
error = cub::DeviceBatchedTopK::MaxKeys(
thrust::raw_pointer_cast(temp_storage.data()) + 1,
temp_storage_bytes,
d_keys_in,
d_keys_out,
segment_sizes,
k_arg,
num_segs,
env);
REQUIRE(error == cudaSuccess);

thrust::sort(keys_out.begin(), keys_out.begin() + k, cuda::std::greater<int>{});
thrust::sort(keys_out.begin() + k, keys_out.begin() + 2 * k, cuda::std::greater<int>{});
REQUIRE(keys_out == thrust::device_vector<int>{8, 7, 6, 9, 8, 7});
}
Loading
Loading