From 9291f7c3b5b04c7e6be0c6a6879f5a2d67dd1c4d Mon Sep 17 00:00:00 2001 From: Bradley Dice Date: Mon, 5 Oct 2026 05:25:36 -0500 Subject: [PATCH] Fix the IVF-PQ list codepacking loops to stride over the whole grid write_list, write_list_flat and run_on_list started at the global row index but advanced by one block's worth of rows, so block b processed every row from its first one to the end of the list. Each row was encoded (or packed/unpacked) up to 16 times, sequentially in the first blocks; with pq_dim 3072 that made an encode launch take ~130 ms. Advance by the whole grid instead, so every row is processed exactly once. The codes are unchanged. --- cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh b/cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh index b1c092c44a..b2d197b3b9 100644 --- a/cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh +++ b/cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -228,7 +228,9 @@ __device__ void run_on_list( uint32_t pq_dim, Action action) { - for (uint32_t ix = threadIdx.x + blockDim.x * blockIdx.x; ix < len; ix += blockDim.x) { + // Grid-stride loop: each vector is processed by exactly one thread. + for (uint32_t ix = threadIdx.x + blockDim.x * blockIdx.x; ix < len; + ix += blockDim.x * gridDim.x) { const uint32_t src_ix = std::holds_alternative(offset_or_indices) ? std::get(offset_or_indices) + ix : std::get(offset_or_indices)[ix]; @@ -248,8 +250,9 @@ __device__ void write_list( Action action) { using subwarp_align = raft::Pow2; - uint32_t stride = subwarp_align::div(blockDim.x); - uint32_t ix = subwarp_align::div(threadIdx.x + blockDim.x * blockIdx.x); + // Grid-stride loop: each vector is processed by exactly one subwarp. + uint32_t stride = subwarp_align::div(blockDim.x) * gridDim.x; + uint32_t ix = subwarp_align::div(threadIdx.x + blockDim.x * blockIdx.x); for (; ix < len; ix += stride) { const uint32_t dst_ix = std::holds_alternative(offset_or_indices) ? std::get(offset_or_indices) + ix @@ -268,8 +271,9 @@ __device__ void write_list_flat(uint8_t* out_flat_codes, Action action) { using subwarp_align = raft::Pow2; - uint32_t stride = subwarp_align::div(blockDim.x); - uint32_t ix = subwarp_align::div(threadIdx.x + blockDim.x * blockIdx.x); + // Grid-stride loop: each vector is processed by exactly one subwarp. + uint32_t stride = subwarp_align::div(blockDim.x) * gridDim.x; + uint32_t ix = subwarp_align::div(threadIdx.x + blockDim.x * blockIdx.x); for (; ix < len; ix += stride) { const uint32_t dst_ix = std::holds_alternative(offset_or_indices) ? std::get(offset_or_indices) + ix