From 154e92f6b40e1e54a412b7ee25f1d514afd33b85 Mon Sep 17 00:00:00 2001 From: sherylll Date: Fri, 2 Oct 2026 02:47:48 -0700 Subject: [PATCH 1/4] Replace u4 wmma with u8 wmma for better portability --- cpp/src/neighbors/detail/nn_descent.cuh | 94 +++++++++++++------------ 1 file changed, 50 insertions(+), 44 deletions(-) diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index ad9ea2e573..164d7a8904 100644 --- a/cpp/src/neighbors/detail/nn_descent.cuh +++ b/cpp/src/neighbors/detail/nn_descent.cuh @@ -1677,6 +1677,21 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) } } +// Masks one nibble half of a packed-4b fragment into `dst`, normalised to 0..15. The high half is +// shifted down so both halves reach the MMA at the same scale and can share one accumulator. +template +__device__ __forceinline__ void select_packed_nibble(const FragT& src, FragT& dst) +{ + constexpr int kWords = FragT::num_elements / 4; + static_assert(FragT::num_elements % 4 == 0, "fragment must be a whole number of u32 words"); + const auto* w = reinterpret_cast(src.x); + auto* d = reinterpret_cast(dst.x); +#pragma unroll + for (int i = 0; i < kWords; ++i) { + d[i] = High ? ((w[i] >> 4) & 0x0f0f0f0fu) : (w[i] & 0x0f0f0f0fu); + } +} + // int4 tensor-core BBQ local join, covering both the symmetric (one quantizer, self-join) and // asymmetric (two quantizers) cases. Modeled on local_join_kernel_wmma: nvcuda::wmma fragments // (u4 x u4 -> s32, shape m8n8k32) replace the popc/dp4a inner product; the accumulator lives in @@ -1714,15 +1729,16 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) cuvs::distance::DistanceType metric, DistEpilogue_t dist_epilogue) { -// int4 sub-byte MMA (nvcuda::wmma experimental::precision::u4) still compiles on every Blackwell -// variant (sm_100/103/110/120/121, verified by disassembly): ptxas lowers it to the same software -// path it's always used since Turing -- unpack each u4 nibble pair into two u8 operands and run -// two native u8 IMMA instructions, summing the partial products. +// The staged bytes are packed 4b; the MMA reads them as u8 and splits the nibbles here rather than +// leaving it to ptxas. u4 MMA is native only through sm_89, and past that ptxas emits this same +// split out of line, once per MMA. Inlining it keeps the IMMAs in the loop body on every arch, and +// the 16x16x16 tile u8 allows drops SUB_PER_DIM to 1. #if (__CUDA_ARCH__ >= 750) using namespace nvcuda; - constexpr int MMA_M = 8; - constexpr int MMA_N = 8; - constexpr int MMA_K = 32; + constexpr int MMA_M = 16; + constexpr int MMA_N = 16; + // A k-step reads 16 packed bytes = 32 codes; each nibble half is a separate MMA over MMA_K. + constexpr int MMA_K = 16; // num_warps = BLOCK_SIZE/32 = 16, arranged as a square WARPS_PER_DIM x WARPS_PER_DIM grid since // 4*4=16 matches exactly; the static_assert is what actually enforces this holds for the // current BLOCK_SIZE, WARPS_PER_DIM itself isn't derived (no trivial constexpr integer sqrt). @@ -1742,19 +1758,13 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) constexpr int SUB_PER_DIM = TILES_PER_DIM / WARPS_PER_DIM; constexpr int WARP_TILE = SUB_PER_DIM * MMA_M; - // Promoted (4-bit-width) staging tile, 128 B/row. Row stride is BBQ_ROW_BYTES + MMA_PAD, not - // just BBQ_ROW_BYTES: sub-byte IMMA loads need at least 16-byte row alignment, and MMA_PAD must - // be a multiple of 16 to preserve that -- but BBQ_ROW_BYTES=128 alone is *also* exactly 32 - // shared-memory banks (4 B/bank), so every row would land on the same bank offset and any - // multi-row access load_matrix_sync does internally would conflict. MMA_PAD=16 breaks that - // exact-32-bank alignment (144 B/row is not a multiple of 128 B) while staying a multiple of 16 - // for the IMMA alignment requirement. constexpr int BBQ_ROW_BYTES = 128; constexpr int MMA_PAD = 16; static_assert(MMA_PAD % 16 == 0, "row padding must preserve 16-byte IMMA row alignment"); - constexpr int ELEMS_PER_TILE = BBQ_ROW_BYTES * 2; // 2 u4 elements/byte - constexpr int K_STEPS_PER_TILE = ELEMS_PER_TILE / MMA_K; - constexpr int ROW_STRIDE_U4 = (BBQ_ROW_BYTES + MMA_PAD) * 2; // row-to-row stride, u4 elements + // Strides are in fragment elements (u8), so plain byte counts, though the data is packed 4b. + constexpr int K_STEPS_PER_TILE = BBQ_ROW_BYTES / MMA_K; + constexpr int KSTEP_BYTES = MMA_K; + constexpr int ROW_STRIDE_U8 = BBQ_ROW_BYTES + MMA_PAD; constexpr int MMA_STORE_STRIDE = SKEWED_MAX_NUM_BI_SAMPLES; @@ -1833,9 +1843,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) [&](const Index_t* col_neighbors, int col_size, auto alias_tag, bool row_resident) { constexpr bool alias_col = decltype(alias_tag)::value; wmma::fragment c_frag[SUB_PER_DIM][SUB_PER_DIM]; -#pragma unroll for (int msub = 0; msub < SUB_PER_DIM; ++msub) { -#pragma unroll for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { wmma::fill_fragment(c_frag[msub][nsub], 0); } @@ -1868,49 +1876,47 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) // per kk and reuse across the other sub-tile index, instead of reloading redundantly inside // a full msub x nsub x kk cross product. // Deliberately not #pragma unroll'd: full unrolling here keeps more fragment live ranges - // simultaneous, driving register pressure up (64/thread, tied with SMEM for the occupancy - // cap) -- letting the compiler pick reduces that at the cost of some intra-warp ILP. + // simultaneous, driving register pressure up -- letting the compiler pick reduces that at + // the cost of some intra-warp ILP. for (int kk = 0; kk < K_STEPS_PER_TILE; ++kk) { - wmma::fragment + wmma::fragment a_frag[SUB_PER_DIM]; - wmma::fragment + wmma::fragment b_frag[SUB_PER_DIM]; const auto* col_buf = alias_col ? s_row_vec : s_col_vec; // compile-time select -#pragma unroll for (int msub = 0; msub < SUB_PER_DIM; ++msub) { const int row0 = warp_id_y * WARP_TILE + msub * MMA_M; - wmma::load_matrix_sync(a_frag[msub], s_row_vec[row0] + kk * (MMA_K / 2), ROW_STRIDE_U4); + wmma::load_matrix_sync(a_frag[msub], s_row_vec[row0] + kk * KSTEP_BYTES, ROW_STRIDE_U8); } -#pragma unroll for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { const int col0 = warp_id_x * WARP_TILE + nsub * MMA_N; - wmma::load_matrix_sync(b_frag[nsub], col_buf[col0] + kk * (MMA_K / 2), ROW_STRIDE_U4); + wmma::load_matrix_sync(b_frag[nsub], col_buf[col0] + kk * KSTEP_BYTES, ROW_STRIDE_U8); } -#pragma unroll - for (int msub = 0; msub < SUB_PER_DIM; ++msub) { -#pragma unroll + // One scratch fragment per operand, reused across the halves. + std::remove_reference_t a_half[SUB_PER_DIM]; + std::remove_reference_t b_half[SUB_PER_DIM]; + auto accumulate_half = [&](auto high_tag) { + constexpr bool High = decltype(high_tag)::value; + for (int msub = 0; msub < SUB_PER_DIM; ++msub) { + select_packed_nibble(a_frag[msub], a_half[msub]); + } for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { - wmma::mma_sync(c_frag[msub][nsub], a_frag[msub], b_frag[nsub], c_frag[msub][nsub]); + select_packed_nibble(b_frag[nsub], b_half[nsub]); } - } + for (int msub = 0; msub < SUB_PER_DIM; ++msub) { + for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { + wmma::mma_sync(c_frag[msub][nsub], a_half[msub], b_half[nsub], c_frag[msub][nsub]); + } + } + }; + accumulate_half(std::false_type{}); + accumulate_half(std::true_type{}); } __syncthreads(); } -#pragma unroll for (int msub = 0; msub < SUB_PER_DIM; ++msub) { const int row0 = warp_id_y * WARP_TILE + msub * MMA_M; -#pragma unroll for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { const int col0 = warp_id_x * WARP_TILE + nsub * MMA_N; wmma::store_matrix_sync( From 3e6a7bbed7e26f5458e1815c437812ca4a14fda3 Mon Sep 17 00:00:00 2001 From: sherylll Date: Fri, 2 Oct 2026 03:20:13 -0700 Subject: [PATCH 2/4] Optimize local_join kernels Skip warps whose output block is out of range; Remove redundant syncthreads in local_join_kernel_wmma Simplify BBQ staging function --- cpp/src/neighbors/detail/nn_descent.cuh | 90 +++++++++++++------------ 1 file changed, 46 insertions(+), 44 deletions(-) diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index 164d7a8904..629d47735b 100644 --- a/cpp/src/neighbors/detail/nn_descent.cuh +++ b/cpp/src/neighbors/detail/nn_descent.cuh @@ -968,6 +968,10 @@ __launch_bounds__(BLOCK_SIZE) wmma::fragment a_frag; wmma::fragment b_frag; wmma::fragment c_frag; + // Lists are usually much shorter than MAX_NUM_BI_SAMPLES, so a warp whose whole 16x16 block + // starts past either extent only produces cells calculate_metric discards -- skip its MMA and + // its store. Staging is cooperative, so this must not skip the barriers. + const bool warp_active = warp_id_y * WMMA_M < list_new_size && warp_id_x * WMMA_N < list_new_size; if (metric != cuvs::distance::DistanceType::BitwiseHamming) { wmma::fill_fragment(c_frag, 0.0); @@ -991,21 +995,23 @@ __launch_bounds__(BLOCK_SIZE) } __syncthreads(); - for (int i = 0; i < TILE_COL_WIDTH / WMMA_K; i++) { + for (int i = 0; warp_active && i < TILE_COL_WIDTH / WMMA_K; i++) { wmma::load_matrix_sync( a_frag, s_nv[warp_id_y * WMMA_M] + i * WMMA_K, TILE_COL_WIDTH + APAD); wmma::load_matrix_sync( b_frag, s_nv[warp_id_x * WMMA_N] + i * WMMA_K, TILE_COL_WIDTH + BPAD); wmma::mma_sync(c_frag, a_frag, b_frag, c_frag); - __syncthreads(); } + __syncthreads(); } - wmma::store_matrix_sync( - s_distances + warp_id_y * WMMA_M * SKEWED_MAX_NUM_BI_SAMPLES + warp_id_x * WMMA_N, - c_frag, - SKEWED_MAX_NUM_BI_SAMPLES, - wmma::mem_row_major); + if (warp_active) { + wmma::store_matrix_sync( + s_distances + warp_id_y * WMMA_M * SKEWED_MAX_NUM_BI_SAMPLES + warp_id_x * WMMA_N, + c_frag, + SKEWED_MAX_NUM_BI_SAMPLES, + wmma::mem_row_major); + } } __syncthreads(); @@ -1034,6 +1040,10 @@ __launch_bounds__(BLOCK_SIZE) __syncthreads(); + // Lists are usually much shorter than MAX_NUM_BI_SAMPLES, so a warp whose whole 16x16 block + // starts past either extent only produces cells calculate_metric discards + const bool warp_active_2 = + warp_id_y * WMMA_M < list_new_size && warp_id_x * WMMA_N < list_old_size; if (metric != cuvs::distance::DistanceType::BitwiseHamming) { wmma::fill_fragment(c_frag, 0.0); for (int step = 0; step < raft::ceildiv(data_dim, TILE_COL_WIDTH); step++) { @@ -1070,21 +1080,23 @@ __launch_bounds__(BLOCK_SIZE) } __syncthreads(); - for (int i = 0; i < TILE_COL_WIDTH / WMMA_K; i++) { + for (int i = 0; warp_active_2 && i < TILE_COL_WIDTH / WMMA_K; i++) { wmma::load_matrix_sync( a_frag, s_nv[warp_id_y * WMMA_M] + i * WMMA_K, TILE_COL_WIDTH + APAD); wmma::load_matrix_sync( b_frag, s_ov[warp_id_x * WMMA_N] + i * WMMA_K, TILE_COL_WIDTH + BPAD); wmma::mma_sync(c_frag, a_frag, b_frag, c_frag); - __syncthreads(); } + __syncthreads(); } - wmma::store_matrix_sync( - s_distances + warp_id_y * WMMA_M * SKEWED_MAX_NUM_BI_SAMPLES + warp_id_x * WMMA_N, - c_frag, - SKEWED_MAX_NUM_BI_SAMPLES, - wmma::mem_row_major); + if (warp_active_2) { + wmma::store_matrix_sync( + s_distances + warp_id_y * WMMA_M * SKEWED_MAX_NUM_BI_SAMPLES + warp_id_x * WMMA_N, + c_frag, + SKEWED_MAX_NUM_BI_SAMPLES, + wmma::mem_row_major); + } __syncthreads(); } @@ -1188,9 +1200,7 @@ __device__ __forceinline__ void stage_promoted_tile(uint8_t (*dst)[RowStride], const Index_t* neighbors, const int count, const int step, - const int native_row_bytes, - const int warp_id, - const int lane_id) + const int native_row_bytes) { static_assert(Layout == bbq_code_layout::packed_1b || Layout == bbq_code_layout::packed_4b, "int4 MMA path supports packed_1b (1b), packed_4b (4b)"); @@ -1201,21 +1211,20 @@ __device__ __forceinline__ void stage_promoted_tile(uint8_t (*dst)[RowStride], constexpr int expansion = Layout == bbq_code_layout::packed_1b ? 4 : 1; constexpr int native_tile = TileBytes / expansion; constexpr int native_tile_u32 = native_tile / 4; - constexpr int num_warps = BLOCK_SIZE / raft::warp_size(); - // packed_4b's native tile is 32 u32s (one per lane). packed_1b's is 8, so a warp-per-row - // map would leave 24 lanes idle; instead pack warp_size/8 = 4 rows into one pass. - static_assert(raft::warp_size() % native_tile_u32 == 0); - constexpr int rows_per_pass = raft::warp_size() / native_tile_u32; - static_assert(MAX_NUM_BI_SAMPLES % (num_warps * rows_per_pass) == 0); + // packed_4b's native tile is 32 u32s (one per thread). packed_1b's is 8, so a warp-per-row + // map would leave 24 lanes idle; instead give each row native_tile_u32 threads. + static_assert(BLOCK_SIZE % native_tile_u32 == 0); + constexpr int rows_per_pass = BLOCK_SIZE / native_tile_u32; + static_assert(MAX_NUM_BI_SAMPLES % rows_per_pass == 0); const int base = step * native_tile; const int remaining = native_row_bytes - base; const int num_load_u32 = (remaining < native_tile ? (remaining > 0 ? remaining : 0) : native_tile) / 4; - for (int i = 0; i < MAX_NUM_BI_SAMPLES / (num_warps * rows_per_pass); ++i) { - const int idx = (i * num_warps + warp_id) * rows_per_pass + lane_id / native_tile_u32; - const int w = lane_id % native_tile_u32; + for (int i = 0; i < MAX_NUM_BI_SAMPLES / rows_per_pass; ++i) { + const int idx = i * rows_per_pass + threadIdx.x / native_tile_u32; + const int w = threadIdx.x % native_tile_u32; if (idx >= count) continue; auto* s = reinterpret_cast(dst[idx]); const auto* src = reinterpret_cast(&quantizer.codes(neighbors[idx], base)); @@ -1849,26 +1858,19 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) } } + // Lists are usually much shorter than MAX_NUM_BI_SAMPLES, so a warp whose whole block starts + // past either extent only produces cells the epilogue discards -- skip its MMA and its store. + // Staging is cooperative, so this must not skip either __syncthreads() in the step loop. + const bool warp_active = warp_id_y * WARP_TILE < new_size && warp_id_x * WARP_TILE < col_size; + for (int step = 0; step < n_tiles; ++step) { if (!row_resident) { - stage_promoted_tile(s_row_vec, - dataset_document, - new_neighbors, - new_size, - step, - doc_row_bytes, - warp_id, - lane_id); + stage_promoted_tile( + s_row_vec, dataset_document, new_neighbors, new_size, step, doc_row_bytes); } if constexpr (!alias_col) { - stage_promoted_tile(s_col_vec, - dataset_query, - col_neighbors, - col_size, - step, - query_row_bytes, - warp_id, - lane_id); + stage_promoted_tile( + s_col_vec, dataset_query, col_neighbors, col_size, step, query_row_bytes); } __syncthreads(); @@ -1878,7 +1880,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) // Deliberately not #pragma unroll'd: full unrolling here keeps more fragment live ranges // simultaneous, driving register pressure up -- letting the compiler pick reduces that at // the cost of some intra-warp ILP. - for (int kk = 0; kk < K_STEPS_PER_TILE; ++kk) { + for (int kk = 0; warp_active && kk < K_STEPS_PER_TILE; ++kk) { wmma::fragment a_frag[SUB_PER_DIM]; wmma::fragment @@ -1915,7 +1917,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) __syncthreads(); } - for (int msub = 0; msub < SUB_PER_DIM; ++msub) { + for (int msub = 0; warp_active && msub < SUB_PER_DIM; ++msub) { const int row0 = warp_id_y * WARP_TILE + msub * MMA_M; for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { const int col0 = warp_id_x * WARP_TILE + nsub * MMA_N; From cd12072b3fcef69cfb1ccb99aa278adb697e097d Mon Sep 17 00:00:00 2001 From: sherylll Date: Wed, 7 Oct 2026 04:14:00 -0700 Subject: [PATCH 3/4] Refactor BBQ local_join kernels: remove redundant template argument --- cpp/src/neighbors/detail/nn_descent.cuh | 52 ++++++++++--------------- 1 file changed, 20 insertions(+), 32 deletions(-) diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index 629d47735b..3886b4afa9 100644 --- a/cpp/src/neighbors/detail/nn_descent.cuh +++ b/cpp/src/neighbors/detail/nn_descent.cuh @@ -1286,7 +1286,6 @@ __device__ __forceinline__ void stage_neighbor_lists(Index_t* new_neighbors, template , @@ -1311,8 +1310,6 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) constexpr int document_planes = cuvs::preprocessing::quantize::bbq::get_code_planes(DocumentLayout); constexpr int query_planes = cuvs::preprocessing::quantize::bbq::get_code_planes(QueryLayout); - static_assert(!SelfJoin || DocumentLayout == QueryLayout, - "a self-join must use the same layout on both operands"); // Both operands are tiled at the same per-plane tile so each step covers the same dimension // range on both sides. QUERY_ROW_BYTES fixes the query row width; the document row width then @@ -1327,6 +1324,9 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) // 1 + 4t 32 32 * 1 = 32 32 32 // 2t + 4t 32 32 * 2 = 64 32 32 // + + constexpr bool SelfJoin = (DocumentLayout == QueryLayout); + constexpr int QUERY_ROW_BYTES = 128; constexpr int BBQ_PAD = alignof(uint32_t); // The document buffer is normally only the A operand (two rows broadcast across a warp), so it @@ -1716,7 +1716,6 @@ __device__ __forceinline__ void select_packed_nibble(const FragT& src, FragT& ds // forced consequence of (MAX_NUM_BI_SAMPLES/MMA_M) / WARPS_PER_DIM, not an arbitrary choice. template , @@ -1849,8 +1848,8 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) // row_resident: the A operand is already staged from a previous phase, which is only true // when the whole row fits one tile (n_tiles == 1) so nothing overwrote it. auto run_phase = - [&](const Index_t* col_neighbors, int col_size, auto alias_tag, bool row_resident) { - constexpr bool alias_col = decltype(alias_tag)::value; + [&](const Index_t* col_neighbors, int col_size, bool row_resident) { + const bool alias_col = (DocumentLayout == QueryLayout) && (col_neighbors == new_neighbors); wmma::fragment c_frag[SUB_PER_DIM][SUB_PER_DIM]; for (int msub = 0; msub < SUB_PER_DIM; ++msub) { for (int nsub = 0; nsub < SUB_PER_DIM; ++nsub) { @@ -1868,7 +1867,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) stage_promoted_tile( s_row_vec, dataset_document, new_neighbors, new_size, step, doc_row_bytes); } - if constexpr (!alias_col) { + if (!alias_col) { stage_promoted_tile( s_col_vec, dataset_query, col_neighbors, col_size, step, query_row_bytes); } @@ -1973,7 +1972,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) __syncthreads(); // ---- Phase 1: new x new ---- - run_phase(new_neighbors, new_size, std::integral_constant{}, false); + run_phase(new_neighbors, new_size, false); for (int step = 0; step < raft::ceildiv(new_size, num_warps); ++step) { const int idx_in_list = step * num_warps + tx / raft::warp_size(); @@ -1994,7 +1993,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) __syncthreads(); // ---- Phase 2: new x old ---- - run_phase(old_neighbors, old_size, std::false_type{}, n_tiles == 1); + run_phase(old_neighbors, old_size, n_tiles == 1); for (int step = 0; step < raft::ceildiv(MAX_NUM_BI_SAMPLES, num_warps); ++step) { const int idx_in_list = step * num_warps + tx / raft::warp_size(); @@ -2535,10 +2534,9 @@ void GNND::local_join( // One launch site for both kernels: they take identical arguments, and the query's layout picks // the path -- packed_4b is the only layout the int4 tensor-core kernel is dispatched for. - auto launch = [&](auto document_layout, auto query_layout, auto self_join_tag) { + auto launch = [&](auto document_layout, auto query_layout) { constexpr auto D = decltype(document_layout)::value; constexpr auto Q = decltype(query_layout)::value; - constexpr bool S = decltype(self_join_tag)::value; auto launch_local_join = [&](auto kernel) { kernel<<>>(graph_.h_graph_new.data_handle(), h_rev_graph_new_.data_handle(), @@ -2562,10 +2560,10 @@ void GNND::local_join( using kernel_id_t = InternalID_t; if constexpr (Q == bbq_code_layout::packed_4b) { launch_local_join( - local_join_kernel_bbq_wmma); + local_join_kernel_bbq_wmma); } else { launch_local_join( - local_join_kernel_bbq_simt); + local_join_kernel_bbq_simt); } }; const bbq_code_layout d = quantizer_document.layout; @@ -2574,47 +2572,38 @@ void GNND::local_join( switch (d) { case bbq_code_layout::packed_1b: launch(std::integral_constant{}, - std::integral_constant{}, - std::true_type{}); + std::integral_constant{}); break; case bbq_code_layout::transposed_2b: launch(std::integral_constant{}, - std::integral_constant{}, - std::true_type{}); + std::integral_constant{}); break; case bbq_code_layout::packed_4b: launch(std::integral_constant{}, - std::integral_constant{}, - std::true_type{}); + std::integral_constant{}); break; case bbq_code_layout::packed_7b: launch(std::integral_constant{}, - std::integral_constant{}, - std::true_type{}); + std::integral_constant{}); break; case bbq_code_layout::packed_8b: launch(std::integral_constant{}, - std::integral_constant{}, - std::true_type{}); + std::integral_constant{}); break; default: RAFT_FAIL("Unsupported BBQ layout for symmetric local join on this branch."); } } else if (d == bbq_code_layout::packed_1b && q == bbq_code_layout::packed_4b) { launch(std::integral_constant{}, - std::integral_constant{}, - std::false_type{}); + std::integral_constant{}); } else if (d == bbq_code_layout::packed_1b && q == bbq_code_layout::transposed_2b) { launch(std::integral_constant{}, - std::integral_constant{}, - std::false_type{}); + std::integral_constant{}); } else if (d == bbq_code_layout::packed_1b && q == bbq_code_layout::transposed_4b) { launch(std::integral_constant{}, - std::integral_constant{}, - std::false_type{}); + std::integral_constant{}); } else if (d == bbq_code_layout::transposed_2b && q == bbq_code_layout::transposed_4b) { launch(std::integral_constant{}, - std::integral_constant{}, - std::false_type{}); + std::integral_constant{}); } else { RAFT_FAIL("Unsupported BBQ layout pair for asymmetric local join."); } @@ -3012,7 +3001,6 @@ void build(raft::resources const& res, if (dataset.has_layout(bbq_code_layout::packed_4b)) { auto kernel = local_join_kernel_bbq_wmma, From bd372afa43d055bbb588d65865a35d76891ea4ed Mon Sep 17 00:00:00 2001 From: sherylll Date: Wed, 7 Oct 2026 04:37:11 -0700 Subject: [PATCH 4/4] Refactor local_join kernels: use compile time degree instead of pass by function argument --- cpp/src/neighbors/detail/nn_descent.cuh | 40 ++++++++------------ cpp/src/neighbors/detail/nn_descent_gnnd.hpp | 2 - 2 files changed, 16 insertions(+), 26 deletions(-) diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index 3886b4afa9..5855464097 100644 --- a/cpp/src/neighbors/detail/nn_descent.cuh +++ b/cpp/src/neighbors/detail/nn_descent.cuh @@ -207,6 +207,8 @@ __device__ __forceinline__ void warp_bitonic_sort(T* element_ptr, const int lane } constexpr int NUM_SAMPLES = 32; +// Width of each row of the device-side kNN graph. +constexpr int DEGREE_ON_DEVICE{32}; // For now, the max. number of samples is 32, so the sample cache size is fixed // to 64 (32 * 2). constexpr int MAX_NUM_BI_SAMPLES = 64; @@ -350,19 +352,16 @@ RAFT_KERNEL add_rev_edges_kernel(const Index_t* graph, } template > -__device__ void insert_to_global_graph(ResultItem elem, - size_t list_id, - ID_t* graph, - DistData_t* dists, - int node_degree, - int* locks) +__device__ void insert_to_global_graph( + ResultItem elem, size_t list_id, ID_t* graph, DistData_t* dists, int* locks) { + constexpr int node_degree = DEGREE_ON_DEVICE; int tx = threadIdx.x; int lane_id = tx % raft::warp_size(); size_t global_idx_base = list_id * node_degree; if (elem.id() == list_id) return; - const int num_segments = raft::ceildiv(node_degree, raft::warp_size()); + constexpr int num_segments = raft::ceildiv(node_degree, raft::warp_size()); int loop_flag = 0; do { @@ -624,7 +623,6 @@ __launch_bounds__(BLOCK_SIZE) const int data_dim, ID_t* graph, DistData_t* dists, - int graph_width, int* locks, DistData_t* l2_norms, cuvs::distance::DistanceType metric, @@ -761,7 +759,7 @@ __launch_bounds__(BLOCK_SIZE) if (idx_in_list >= list_new_size) continue; auto min_elem = get_min_item(s_list[idx_in_list], idx_in_list, new_neighbors, s_distances); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -858,7 +856,7 @@ __launch_bounds__(BLOCK_SIZE) } if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } #endif @@ -889,7 +887,6 @@ __launch_bounds__(BLOCK_SIZE) const int data_dim, ID_t* graph, DistData_t* dists, - int graph_width, int* locks, DistData_t* l2_norms, cuvs::distance::DistanceType metric, @@ -1032,7 +1029,7 @@ __launch_bounds__(BLOCK_SIZE) if (idx_in_list >= list_new_size) continue; auto min_elem = get_min_item(s_list[idx_in_list], idx_in_list, new_neighbors, s_distances); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -1131,7 +1128,7 @@ __launch_bounds__(BLOCK_SIZE) } if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } #endif @@ -1302,7 +1299,6 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) quantizer_view dataset_query, ID_t* graph, DistData_t* dists, - int graph_width, int* locks, cuvs::distance::DistanceType metric, DistEpilogue_t dist_epilogue) @@ -1549,7 +1545,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) SKEWED_MAX_NUM_BI_SAMPLES, new_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -1665,7 +1661,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) SKEWED_MAX_NUM_BI_SAMPLES, old_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -1681,7 +1677,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) SKEWED_MAX_NUM_BI_SAMPLES, new_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[list_idx], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[list_idx], graph, dists, locks); } } } @@ -1732,7 +1728,6 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) const quantizer_view dataset_query, ID_t* graph, DistData_t* dists, - int graph_width, int* locks, cuvs::distance::DistanceType metric, DistEpilogue_t dist_epilogue) @@ -1985,7 +1980,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) MMA_STORE_STRIDE, new_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -2006,7 +2001,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) MMA_STORE_STRIDE, old_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[idx_in_list], graph, dists, locks); } } @@ -2017,7 +2012,7 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) auto min_elem = get_min_item( s_list[list_idx], idx_in_list, new_neighbors, s_distances, false, MMA_STORE_STRIDE, new_size); if (min_elem.id() < gridDim.x) { - insert_to_global_graph(min_elem, s_list[list_idx], graph, dists, graph_width, locks); + insert_to_global_graph(min_elem, s_list[list_idx], graph, dists, locks); } } #endif // (__CUDA_ARCH__ >= 750) @@ -2444,7 +2439,6 @@ void GNND::local_join(cudaStream_t stream, DistEpilogue_t dist_ ndim_, graph_buffer_.data_handle(), dists_buffer_.data_handle(), - DEGREE_ON_DEVICE, d_locks_.data_handle(), l2_norms_.data_handle(), build_config_.metric, @@ -2461,7 +2455,6 @@ void GNND::local_join(cudaStream_t stream, DistEpilogue_t dist_ ndim_, graph_buffer_.data_handle(), dists_buffer_.data_handle(), - DEGREE_ON_DEVICE, d_locks_.data_handle(), l2_norms_.data_handle(), build_config_.metric, @@ -2549,7 +2542,6 @@ void GNND::local_join( quantizer_query, graph_buffer_.data_handle(), dists_buffer_.data_handle(), - DEGREE_ON_DEVICE, d_locks_.data_handle(), build_config_.metric, dist_epilogue); diff --git a/cpp/src/neighbors/detail/nn_descent_gnnd.hpp b/cpp/src/neighbors/detail/nn_descent_gnnd.hpp index 07db8b7064..095b08865f 100644 --- a/cpp/src/neighbors/detail/nn_descent_gnnd.hpp +++ b/cpp/src/neighbors/detail/nn_descent_gnnd.hpp @@ -21,8 +21,6 @@ namespace cuvs::neighbors::nn_descent::detail { using DistData_t = float; -constexpr int DEGREE_ON_DEVICE{32}; -constexpr int SEGMENT_SIZE{32}; constexpr int counter_interval{100}; template struct InternalID_t;