77from dataclasses import dataclass
88from pathlib import Path
99
10- import lucene
1110
12-
13- REPO_ROOT = Path (__file__ ).resolve ().parents [2 ]
11+ REPO_ROOT = Path (__file__ ).resolve ().parents [3 ]
1412HNSW_CODEC = "Lucene101AcceleratedHNSWCodec"
1513CAGRA_HNSW_BASE_LAYER_CODEC = "Lucene101AcceleratedHNSWBaseLayerCodec"
1614CAGRA_HNSW_MULTI_LAYER_CODEC = "Lucene101AcceleratedHNSWMultiLayerCodec"
2826 "Lucene101AcceleratedHNSWBinaryQuantizedCodec" ,
2927 "Lucene101AcceleratedHNSWScalarQuantizedCodec" ,
3028)
31- BASIC_GPU_CASES = ("hnsw" , "cagra" , "hnsw-single" , "cagra-single" )
29+ CAGRA_SEARCH_CASES = (
30+ "cagra-search-default" ,
31+ "cagra-search-single" ,
32+ "cagra-search-1seg" ,
33+ "cagra-search-10seg" ,
34+ "cagra-search-10seg-force-1" ,
35+ "cagra-search-100seg-force-10" ,
36+ )
37+ BASIC_GPU_CASES = (
38+ "hnsw" ,
39+ "cagra-search-default" ,
40+ "hnsw-single" ,
41+ "cagra-search-single" ,
42+ )
3243SEGMENT_GPU_CASES = (
3344 "hnsw-1seg" ,
34- "cagra-1seg" ,
45+ "cagra-search- 1seg" ,
3546 "hnsw-10seg" ,
36- "cagra-10seg" ,
47+ "cagra-search- 10seg" ,
3748 "hnsw-10seg-force-1" ,
38- "cagra-10seg-force-1" ,
49+ "cagra-search- 10seg-force-1" ,
3950 "hnsw-100seg-force-10" ,
40- "cagra-100seg-force-10" ,
51+ "cagra-search- 100seg-force-10" ,
4152)
4253CPU_HNSW_CASES = (
4354 "hnsw-cpu" ,
5566 "hnsw-cpu" ,
5667 "cagra-hnsw-1layer" ,
5768 "cagra-hnsw-3layer" ,
58- "cagra" ,
69+ "cagra-search-default " ,
5970)
6071CASE_GROUPS = {
6172 "gpu" : BASIC_GPU_CASES ,
6273 "gpu-basic" : BASIC_GPU_CASES ,
6374 "gpu-segments" : SEGMENT_GPU_CASES ,
6475 "segments" : SEGMENT_GPU_CASES ,
76+ "cagra-search" : CAGRA_SEARCH_CASES ,
77+ "gpu-cagra-search" : CAGRA_SEARCH_CASES ,
6578 "cpu-hnsw" : CPU_HNSW_CASES ,
6679 "cagra-hnsw" : CAGRA_HNSW_CASES ,
6780 "algorithm-matrix" : ALGORITHM_MATRIX_CASES ,
@@ -206,6 +219,8 @@ def fvec(jarray, values):
206219
207220
208221def init_vm (cuvs_java_jar , cuvs_lucene_jar ):
222+ import lucene
223+
209224 java_library_path = os .environ .get ("JAVA_LIBRARY_PATH" ) or os .environ .get (
210225 "LD_LIBRARY_PATH"
211226 )
@@ -385,9 +400,9 @@ def build_named_case(name):
385400 segment_count = 3 ,
386401 force_merge_target = (1 if force_merge else 0 ),
387402 )
388- if name == "cagra" :
403+ if name in { "cagra" , "cagra-search" , "cagra-search-default" } :
389404 return matrix_case (
390- name = "cagra" ,
405+ name = "cagra-search-default " ,
391406 codec_name = CAGRA_CODEC ,
392407 expected_suffixes = (".vcag" , ".vemc" ),
393408 require_cuvs = True ,
@@ -405,9 +420,9 @@ def build_named_case(name):
405420 expected_suffixes = (".vex" , ".vem" ),
406421 require_cuvs = require_cuvs ,
407422 )
408- if name in {"cagra-single" , "single-cagra" }:
423+ if name in {"cagra-single" , "single-cagra" , "cagra-search-single" }:
409424 return SmokeCase (
410- name = "cagra-single" ,
425+ name = "cagra-search- single" ,
411426 codec_name = CAGRA_CODEC ,
412427 row_count = 1 ,
413428 dims = matrix_dims ,
@@ -417,20 +432,33 @@ def build_named_case(name):
417432 )
418433 if name == "hnsw-1seg" :
419434 return build_segment_case (name , HNSW_CODEC , (".vex" , ".vem" ), require_cuvs , 1 , 0 )
420- if name == "cagra-1seg" :
421- return build_segment_case (name , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 1 , 0 )
435+ if name in {"cagra-1seg" , "cagra-search-1seg" }:
436+ return build_segment_case (
437+ "cagra-search-1seg" , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 1 , 0
438+ )
422439 if name == "hnsw-10seg" :
423440 return build_segment_case (name , HNSW_CODEC , (".vex" , ".vem" ), require_cuvs , 10 , 0 )
424- if name == "cagra-10seg" :
425- return build_segment_case (name , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 10 , 0 )
441+ if name in {"cagra-10seg" , "cagra-search-10seg" }:
442+ return build_segment_case (
443+ "cagra-search-10seg" , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 10 , 0
444+ )
426445 if name == "hnsw-10seg-force-1" :
427446 return build_segment_case (name , HNSW_CODEC , (".vex" , ".vem" ), require_cuvs , 10 , 1 )
428- if name == "cagra-10seg-force-1" :
429- return build_segment_case (name , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 10 , 1 )
447+ if name in {"cagra-10seg-force-1" , "cagra-search-10seg-force-1" }:
448+ return build_segment_case (
449+ "cagra-search-10seg-force-1" , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 10 , 1
450+ )
430451 if name == "hnsw-100seg-force-10" :
431452 return build_segment_case (name , HNSW_CODEC , (".vex" , ".vem" ), require_cuvs , 100 , 10 )
432- if name == "cagra-100seg-force-10" :
433- return build_segment_case (name , CAGRA_CODEC , (".vcag" , ".vemc" ), True , 100 , 10 )
453+ if name in {"cagra-100seg-force-10" , "cagra-search-100seg-force-10" }:
454+ return build_segment_case (
455+ "cagra-search-100seg-force-10" ,
456+ CAGRA_CODEC ,
457+ (".vcag" , ".vemc" ),
458+ True ,
459+ 100 ,
460+ 10 ,
461+ )
434462 if name in {"cagra-hnsw-1layer" , "cagra-hnsw-base" , "cagra-hnsw-base-layer" }:
435463 return matrix_case (
436464 name = "cagra-hnsw-1layer" ,
@@ -767,22 +795,54 @@ def hit_ids(stored_fields, hits):
767795 return [stored_fields .document (hit .doc ).get (ID_FIELD ) for hit in hits ]
768796
769797
770- def assert_search_results (searcher , jarray , case , query_ids , inactive_ids ):
798+ def squared_l2 (left , right ):
799+ return sum ((a - b ) * (a - b ) for a , b in zip (left , right ))
800+
801+
802+ def exact_neighbor_doc_names (query_id , active_vector_ids , dims , limit ):
803+ query_vector = vector_for (query_id , dims )
804+ ranked = sorted (
805+ active_vector_ids ,
806+ key = lambda doc_id : (squared_l2 (query_vector , vector_for (doc_id , dims )), doc_id ),
807+ )
808+ return [f"doc-{ doc_id } " for doc_id in ranked [:limit ]]
809+
810+
811+ def assert_search_results (searcher , jarray , case , query_ids , active_vector_ids , inactive_ids ):
771812 from org .apache .lucene .search import KnnFloatVectorQuery
772813
773814 stored_fields = searcher .storedFields ()
774- top_k = min (case .top_k , max (1 , case . row_count - len (inactive_ids )))
815+ top_k = min (case .top_k , max (1 , len (active_vector_ids )))
775816 for query_id in query_ids :
776817 query = KnnFloatVectorQuery (
777818 VECTOR_FIELD , fvec (jarray , vector_for (query_id , case .dims )), top_k
778819 )
779820 ids = hit_ids (stored_fields , searcher .search (query , top_k ).scoreDocs )
780821 expected = f"doc-{ query_id } "
781- if expected not in ids :
782- raise AssertionError (f"{ case .name } : expected { expected } in top { top_k } , got { ids } " )
822+ if len (ids ) != top_k :
823+ raise AssertionError (
824+ f"{ case .name } : expected { top_k } search results, got { len (ids )} : { ids } "
825+ )
826+ if len (ids ) != len (set (ids )):
827+ raise AssertionError (f"{ case .name } : duplicate docs returned: { ids } " )
828+ if ids [0 ] != expected :
829+ raise AssertionError (f"{ case .name } : expected { expected } at rank 1, got { ids } " )
783830 bad_ids = [doc_id for doc_id in ids if doc_id in inactive_ids ]
784831 if bad_ids :
785832 raise AssertionError (f"{ case .name } : inactive docs returned: { bad_ids } " )
833+ expected_candidates = set (
834+ exact_neighbor_doc_names (
835+ query_id , active_vector_ids , case .dims , min (len (active_vector_ids ), top_k * 2 + 10 )
836+ )
837+ )
838+ unexpected_ids = [
839+ doc_id for doc_id in ids [: min (5 , len (ids ))] if doc_id not in expected_candidates
840+ ]
841+ if unexpected_ids :
842+ raise AssertionError (
843+ f"{ case .name } : nearest hits { unexpected_ids } were outside the exact "
844+ f"top { len (expected_candidates )} candidates for { expected } "
845+ )
786846
787847
788848def assert_filtered_search (searcher , jarray , case , query_id ):
@@ -906,7 +966,9 @@ def run_case(case, codec_class, codec_cache, jarray):
906966 assert_vector_metadata (reader , VECTOR_FIELD , expected_vector_count , case .dims )
907967 assert_segment_topology (reader , case )
908968 searcher = IndexSearcher (reader )
909- assert_search_results (searcher , jarray , case , query_ids , inactive_doc_names )
969+ assert_search_results (
970+ searcher , jarray , case , query_ids , active_vector_ids , inactive_doc_names
971+ )
910972 assert_filtered_search (searcher , jarray , case , query_ids [0 ])
911973 telemetry = writer_telemetry (codec )
912974 writer_path = observed_writer_path (case , telemetry )
0 commit comments