Skip to content
Merged
Changes from all commits
Commits
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
16 changes: 10 additions & 6 deletions cpp/src/neighbors/ivf_pq/ivf_pq_codepacking.cuh
Original file line number Diff line number Diff line change
@@ -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
*/

Expand Down Expand Up @@ -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<uint32_t>(offset_or_indices)
? std::get<uint32_t>(offset_or_indices) + ix
: std::get<const uint32_t*>(offset_or_indices)[ix];
Expand All @@ -248,8 +250,9 @@ __device__ void write_list(
Action action)
{
using subwarp_align = raft::Pow2<SubWarpSize>;
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<uint32_t>(offset_or_indices)
? std::get<uint32_t>(offset_or_indices) + ix
Expand All @@ -268,8 +271,9 @@ __device__ void write_list_flat(uint8_t* out_flat_codes,
Action action)
{
using subwarp_align = raft::Pow2<SubWarpSize>;
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<uint32_t>(offset_or_indices)
? std::get<uint32_t>(offset_or_indices) + ix
Expand Down
Loading