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