diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index ad9ea2e573..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, @@ -968,6 +965,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 +992,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(); @@ -1026,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); } } @@ -1034,6 +1037,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 +1077,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(); } @@ -1119,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 @@ -1188,9 +1197,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 +1208,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)); @@ -1277,7 +1283,6 @@ __device__ __forceinline__ void stage_neighbor_lists(Index_t* new_neighbors, template , @@ -1294,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) @@ -1302,8 +1306,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 @@ -1318,6 +1320,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 @@ -1540,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); } } @@ -1656,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); } } @@ -1672,11 +1677,26 @@ 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); } } } +// 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 @@ -1692,7 +1712,6 @@ RAFT_KERNEL __launch_bounds__(BLOCK_SIZE) // forced consequence of (MAX_NUM_BI_SAMPLES/MMA_M) / WARPS_PER_DIM, not an arbitrary choice. template , @@ -1709,20 +1728,20 @@ 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) { -// 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 +1761,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; @@ -1830,37 +1843,28 @@ 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]; -#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); } } + // 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); + if (!alias_col) { + stage_promoted_tile( + s_col_vec, dataset_query, col_neighbors, col_size, step, query_row_bytes); } __syncthreads(); @@ -1868,49 +1872,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. - for (int kk = 0; kk < K_STEPS_PER_TILE; ++kk) { - wmma::fragment + // simultaneous, driving register pressure up -- letting the compiler pick reduces that at + // the cost of some intra-warp ILP. + for (int kk = 0; warp_active && kk < K_STEPS_PER_TILE; ++kk) { + 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) { + for (int msub = 0; warp_active && 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( @@ -1965,7 +1967,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(); @@ -1978,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); } } @@ -1986,7 +1988,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(); @@ -1999,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); } } @@ -2010,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) @@ -2437,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, @@ -2454,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, @@ -2527,10 +2527,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(), @@ -2543,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); @@ -2554,10 +2552,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; @@ -2566,47 +2564,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."); } @@ -3004,7 +2993,6 @@ void build(raft::resources const& res, if (dataset.has_layout(bbq_code_layout::packed_4b)) { auto kernel = local_join_kernel_bbq_wmma, 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;