Skip to content

Vamana build optimization set and fp16 support - #2264

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

Vamana build optimization set and fp16 support#2264
bkarsin wants to merge 30 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

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?

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.

Having just lane 0 write it would require a sync after. Seems simpler/cleaner to just have all threads write, since they have the same value.

Comment thread cpp/src/neighbors/detail/vamana/vamana_structs.cuh Outdated
__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.

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.

I can change this throughout, but I can't find the docs in raft on the different INLINE macros.

Comment thread cpp/src/neighbors/detail/vamana/priority_queue.cuh Outdated
Comment thread cpp/src/neighbors/detail/vamana/vamana_build.cuh Outdated
__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.

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 scatter kernel is not so simple. I could do it but would need to re-work the code a bit more to use a strided mdspan or something. Can work on this if it's a blocker for the PR.

Comment thread cpp/src/neighbors/detail/vamana/vamana_structs.cuh Outdated
Comment thread cpp/src/neighbors/detail/vamana/vamana_build.cuh Outdated
Comment thread cpp/src/neighbors/detail/vamana/robust_prune.cuh
// 1024 threads/SM, so 12 * 128 = 1536 exceeds the limit and ptxas emits
// ".minnctapersm out of range ... will be ignored". Use 8 (8 * 128 = 1024) on those
// architectures and keep the tuned value of 12 elsewhere (>=1536 threads/SM).
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 720 || __CUDA_ARCH__ == 750)

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.

This macro can be used only in device code, right? For host code we can use raft::resource::get_device_properties. Fore example we have this elsewhere:

/**
 * @brief Returns true if the fused distance NN implementation should be used.
 *
 * On Ampere (SM <= 8.x) always use fused.
 * On Hopper (SM 9.x) use fused when m or n >= 4096.
 * On Blackwell (SM >= 10.x) use unfused.
 */
template <typename MathT, typename IdxT, typename LabelT>
bool use_fused(const raft::resources& handle, IdxT m, IdxT n, IdxT k)
{
  cudaDeviceProp prop;
  prop = raft::resource::get_device_properties(handle);
  if (prop.major <= 8) {
    // Use fused for Ampere or before
    return true;
  } else if (prop.major == 9 && (m >= 4096 || n >= 4096)) {
    // On Hopper if m, n are bigger than 4096, use fused
    return true;
  } else if (prop.major >= 10) {
    // On Blackwell onwards, use unfused
    return false;
  }
  return false;
}

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 only used here to define the launch_bounds, since older hardware has a smaller max threads per SM, and using the smaller value on new hardware hurts performance. It's a compile-time parameter so it needs to be a macro, I believe. It looks like the use_fused example determines which kernel to use during runtime, though.

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