Skip to content

Vamana build optimization set and fp16 support#2264

Open
bkarsin wants to merge 28 commits into
NVIDIA:mainfrom
bkarsin:vamana-build-opt
Open

Vamana build optimization set and fp16 support#2264
bkarsin wants to merge 28 commits into
NVIDIA:mainfrom
bkarsin:vamana-build-opt

Conversation

@bkarsin

@bkarsin bkarsin commented Jun 25, 2026

Copy link
Copy Markdown
Contributor

Series of GPU Vamana build performance optimizations that addresses #2178 and #1757. Initial estimates from #1757 were not accurate, so many other optimizations were tried (some abandoned, some successful). This PR includes:

GreedySearch optimizations:

  • Multi-warp blocks - increases occupancy of GreedySearch
  • Reduce shared memory used per warp
  • fp16 query approximation (along with fp16 support)

RobustPrune optimizations:

  • Multi-warp block support (block size depends on graph degree)
  • Caching accepted candidate vectors during occlusion loop
  • Remove redundant merge of candidate lists and avoid syncs

General optimizations:

  • Replace prefix_sums kernel with cub variant
  • Re-use distances from GreedySearch in RobustPrune kernel
  • fp16 support added

Together these optimizations give significant speedups across all configs with minimal recall variance compared to the current baseline. I benchmarked performance across a range of synthetic datasets and two real-world datasets. (NOTE: finishing benchmarks and will update tables below once they are all collected).

Synthetic dataset build tests (all 1M vector datasets)

    RTX pro 6000 RTX pro 6000 H100 H100 L4 L4
TYPE Dim/Degree BASELINE OPT BASELINE OPT BASELINE OPT
fp32 64/32 1.11 1.07 1.35 1.19 9.83 3.57
fp32 768 / 32 5.98 4.03 6.01 4.92 36.6 25.8
fp32 960 / 32 7.82 5.61 8.23 6.52 50.1 38.7
fp32 64/64 4.4 4.05 6.14 5.4 16.4 14.8
fp32 768/64 28.8 18.2 33.7 25.3 194 125
fp32 960/64 39.4 24.9 50.7 36.5 251 185
fp16 64/32   1.05       3.8
fp16 768 / 32   2.5       11.1
fp16 960 / 32   3.13       14.1
fp16 64/64   4.01       13.6
fp16 768/64   10.4       53.8
fp16 960/64   14.2       65.6
int8 64/32 1.06 0.966 1.69 1.74 3.5 3.35
int8 768 / 32 2.98 2.38 4.82 3.69 11.2 9.21
int8 960 / 32 3.85 3 6.39 4.32 15.1 12
int8 64/64 4.15 3.61 6.36 6.05 14.4 13.3
int8 768/64 12.9 9.74 22.4 14.9 51.7 39.7
int8 960/64 17.1 12.5 28.8 19.9 72.1 52.1

Also tested real-world BIGANN 10M (uint8 128D) and GIST (fp32 960D) datasets:

    RTX pro 6000 RTX pro 6000 H100 H100 L4 L4
Dataset deg / iters BASELINE OPT BASELINE OPT BASELINE OPT
BIGANN 32 / 1.0 12.017 10.736 13.1152 12.58 40.3484 39.887
BIGANN 32 / 2.0 25.942 23.229 32.0109 30.4 90.6138 89.994
BIGANN 64 / 1.0 36.642 30.688 48.1399 43.65 126.248 120.394
BIGANN 64 / 2.0 86.221 72.026 118.444 107.91 299.2 285.543
GIST (fp32) 32 / 1.0 6.39807 4.9804 6.773 5.51 41.48 32.3789
GIST (fp32) 32 / 2.0 16.2624 12.1026 19.0785 13.7675 103.197 76.84
GIST (fp32) 64 / 1.0 22.9393 16.74 27.2307 20.1767 150.2 109.98
GIST (fp32) 64 / 2.0 60.1051 41.452 71.1611 53.6482 380.185 262.1
GIST (fp16) 32 / 1.0   2.597   3.80264    
GIST (fp16) 32 / 2.0   6.176   8.692    
GIST (fp16) 64 / 1.0   9.965   13.793    
GIST (fp16) 64 / 2.0   23.95   34.1922    

bkarsin added 15 commits June 15, 2026 14:11
…e vector in shared memory in the RobustPrune occlusion loop

(cherry picked from commit f45fd1b49283434eb4a3017da069ead501e938c3)
…efix sum with cub scan and hoist per-batch reverse-edge allocations

(cherry picked from commit 3b8650f5ebea52421f187332a2f6f3bdd599c42e)
…lusion across multiple warps per query (raise occupancy)

(cherry picked from commit 2e02f938f97e65ca073daf07397d538432b52867)
… query->existing-edge distances in the RobustPrune merge (avoid recompute)

(cherry picked from commit 041d355c585b7f98302d054b80ac287d637a07e4)
…oords in FP16 smem for dim>=512 to raise GreedySearch occupancy (salvage of N8)
…nce (one warp) instead of redundantly on all 128 threads, then broadcast
…ck (4 vs 8) to raise occupancy/MLP on the degree-64 occlusion sweep
@copy-pr-bot

copy-pr-bot Bot commented Jun 25, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@cjnolet cjnolet added improvement Improves an existing functionality non-breaking Introduces a non-breaking change labels Jul 1, 2026
@cjnolet cjnolet moved this to In Progress in Unstructured Data Processing Jul 1, 2026
#define KERNEL_TIMING (RAFT_LOG_ACTIVE_LEVEL <= RAPIDS_LOGGER_LOG_LEVEL_DEBUG)

template <typename accT, typename IdxT>
__global__ void gather_query_sizes(QueryCandidates<IdxT, accT>* query_list,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you have a sense for how much this is adding to the binary size? @divyegala should be able to explain how to see the deployed metrics for PRs.

@cjnolet
cjnolet requested review from a team as code owners July 15, 2026 15:04
@jamxia155

Copy link
Copy Markdown
Contributor

/ok to test c055c60

@jamxia155

Copy link
Copy Markdown
Contributor

/ok to test df13b42

@pytest.mark.parametrize("dtype", [np.float32, np.int8, np.uint8])
@pytest.mark.parametrize("dtype", [np.float32, np.float16, np.int8, np.uint8])
def test_vamana_build_basic(dtype):
n_rows, n_cols = 1000, 16

@tarang-jain tarang-jain Jul 22, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we add at least one test at larger dims? The current new FP16 coverage uses only 16 dimensions, so I think we dont test greedy_search_use_fp16_query_smem()?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is just for the python test. Added fp16 cpp tests (cpp/tests/neighbors/ann_vamana/test_half_uint32_t.cu) that cover a wide range of params.

@jamxia155

Copy link
Copy Markdown
Contributor

/ok to test 350ed89

@copy-pr-bot

copy-pr-bot Bot commented Jul 22, 2026

Copy link
Copy Markdown

/ok to test 350ed89

@jamxia155, there was an error processing your request: E2

See the following link for more information: https://docs.gha-runners.nvidia.com/cpr/e/2/

@jamxia155

Copy link
Copy Markdown
Contributor

/ok to test 74caba9

@jamxia155

Copy link
Copy Markdown
Contributor

/ok to test b8238f3

Comment thread cpp/tests/neighbors/ann_vamana.cuh Outdated
Comment on lines +258 to +264
} else if constexpr (std::is_same_v<DataT, half>) {
// raft::random::normal requires std::is_floating_point (excludes half); use uniform,
// which supports half. uniformInt (the integer branch below) rejects half.
raft::random::uniform(
handle_, r, database.data(), ps.n_rows * ps.dim, DataT(0.1), DataT(2.0));
raft::random::uniform(
handle_, r, search_queries.data(), ps.n_queries * ps.dim, DataT(0.1), DataT(2.0));

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I suggest that we just generate normal data and then cast that to fp16 instead of generating fp16 in uniform distribution

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed


num_neighbors = degree;
__syncthreads();
num_neighbors[warpIdx] = degree;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

maybe make only lane 0 write to this?

}

// Checks current list to see if a node as previously been visited
__inline__ __device__ bool check_visited(IdxT target, accT dist)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I believe this function is no longer used anywhere anymore. Should we delete it?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Removed

__float2half(0.0f), __float2half(0.0f), __float2half(0.0f), __float2half(0.0f)};
for (int i = threadIdx.x; i < src_vec->Dim; i += 4 * blockDim.x) {
temp_dst[0] = dst_vec->coords[i];
if (i + 32 < src_vec->Dim) temp_dst[1] = dst_vec->coords[i + 32];

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we use blockDim.x instead of hardcoding to 32 here?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I agree with this. We should not be hardcoding things because the kernel invocation can change (or the same kernel may be called somewhere else in the future).

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is run by a single warp within the block, and uses a warp width stride. I can change it to a define macro if that helps?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Replaced with raft::WarpSize

@tarang-jain tarang-jain left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

There are a lot of changes in this PR where I want to basically tell the exact same things -- using raft macros and raft primtives wherever possible. We should certainly avoid writing small kernels which can easily be a lambda for raft::linalg. It helps keeps things readable and binary size increases can become slightly more predictable. For example, earlier we have found that raft::linalg::map kernels generally add less to the binary size than thrust::for_each (or similar thrust primitives).

Node<accT>* topk_pq_mem)
{
int n = dataset.extent(0);
const int warpIdx = threadIdx.x / 32;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

here as well. Lets not hardcode 32. Lets use raft::WarpSize instead.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Replaced all warp width values with raft::WarpSize throughout

Node<SUMTYPE> input_data,
SUMTYPE* cur_max_val,
int* max_idx)
__inline__ __device__ void parallel_pq_max_enqueue_warp(Node<SUMTYPE>* pq,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nit: we have macros to mark kernels: such as _RAFT_INLINE

@tarang-jain tarang-jain Jul 23, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wouldnt say these macros should block the PR from merging.

}

if (threadIdx.x == 31) {
if (laneId == 31) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Again I am marking more places where things are hardcoded. We have raft variables to take care of this.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Replace with raft::WarpSize-1

#define KERNEL_TIMING (RAFT_LOG_ACTIVE_LEVEL <= RAPIDS_LOGGER_LOG_LEVEL_DEBUG)

template <typename accT, typename IdxT>
__global__ void gather_query_sizes(QueryCandidates<IdxT, accT>* query_list,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I would argue that we can achieve these things with raft linalg primitives (such as raft::linalg::map rather than having to write a new kernel).

__global__ void scatter_prefix_offsets(QueryCandidates<IdxT, accT>* query_list,
const int* edge_offsets,
int count)
{

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same here: we should be able to achieve these with raft primitives.

}

template <typename T>
__host__ __device__ inline int greedy_search_query_smem_elem_size(int dim)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Again, lets use our raft macros: RAFT_HOST_DEVICE_INLINE

// Find which node IDs have reverse edges and their indices in the reverse edge list
thrust::device_vector<IdxT> edge_dest_vec(edge_dest.data_handle(),
edge_dest.data_handle() + total_edges);
RAFT_CUDA_TRY(cudaMemcpyAsync(thrust::raw_pointer_cast(edge_dest_vec.data()),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

raft::copy please!

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

int accept_count = 0;
if (s_res_size > degree) {
if (threadIdx.x == 0) s_accept_count = 0;
__syncthreads();

@tarang-jain tarang-jain Jul 23, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

do we need to sync the whole block here? We do it already at line 234. I see that s_accept_count is set to zero and you might need the whole block to see this before moving forward, but how about all threads set that variable rather than just thread zero setting it? Let me know if my understanding is not correct here.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All the threads in the block use the shared s_accept_count after this, and any thread may increment it. If we have all threads set it to 0 and no sync, we may lose later increments by a trailing thread/warp re-setting it to 0.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

improvement Improves an existing functionality non-breaking Introduces a non-breaking change

Projects

Status: In Progress

Development

Successfully merging this pull request may close these issues.

6 participants