From 6fd5b2fc1f13559b98ffa9f1e57c0e41040edcdc Mon Sep 17 00:00:00 2001 From: zhenghaoz Date: Wed, 12 Aug 2026 13:25:10 +0800 Subject: [PATCH] perf(hnsw): reuse visited workspaces --- internal/core/group_by_hnsw.go | 11 ++-- internal/core/hnsw.go | 42 +++++++++------- internal/core/hnsw_parallel.go | 17 ++++--- internal/core/hnsw_quantized.go | 34 +++++++------ internal/core/hnsw_rabitq.go | 20 +++++--- internal/core/hnsw_sparse.go | 42 +++++++++------- internal/core/hnsw_test.go | 1 + internal/core/hnsw_visited.go | 51 +++++++++++++++++++ internal/core/hnsw_visited_test.go | 80 ++++++++++++++++++++++++++++++ 9 files changed, 228 insertions(+), 70 deletions(-) create mode 100644 internal/core/hnsw_visited.go create mode 100644 internal/core/hnsw_visited_test.go diff --git a/internal/core/group_by_hnsw.go b/internal/core/group_by_hnsw.go index ed50331..d6ba158 100644 --- a/internal/core/group_by_hnsw.go +++ b/internal/core/group_by_hnsw.go @@ -66,6 +66,7 @@ func expandHNSWGroups( publicScore func(score float32) float32, nodeBetter func(left, right hnswScoredNode) bool, prefetch func(neighbors []int), + visited *hnswVisited, ) ([]GroupResult, error) { accumulator := newGroupAccumulator(metric, options.TopKPerGroup) groups := make(map[string]struct{}, min(options.GroupCount, len(initial))) @@ -93,12 +94,12 @@ func expandHNSWGroups( } frontier := ailego.NewHeap(nodeBetter) - visited := make([]bool, len(keys)) + visited.reset(len(keys)) for _, node := range initial { - if node.position < 0 || node.position >= len(keys) || visited[node.position] { + if node.position < 0 || node.position >= len(keys) || visited.seen(node.position) { continue } - visited[node.position] = true + visited.mark(node.position) frontier.Push(node) } for frontier.Len() != 0 && len(groups) < options.GroupCount { @@ -111,10 +112,10 @@ func expandHNSWGroups( prefetch(adjacent) } for _, neighbor := range adjacent { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := scoreAt(neighbor) if err != nil { return nil, fmt.Errorf("core: score HNSW group expansion node %d: %w", neighbor, err) diff --git a/internal/core/hnsw.go b/internal/core/hnsw.go index 100dad5..de0a479 100644 --- a/internal/core/hnsw.go +++ b/internal/core/hnsw.go @@ -355,6 +355,8 @@ func (i *HNSWIndex) Neighbors(key uint64, level int) ([]uint64, error) { } func (i *HNSWIndex) insertBuiltNode(ctx context.Context, position int) error { + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) level := i.levels[position] if i.entryPoint < 0 { i.entryPoint = position @@ -364,7 +366,7 @@ func (i *HNSWIndex) insertBuiltNode(ctx context.Context, position int) error { entry := i.entryPoint query := i.vectorAt(position) for currentLevel := i.maxLevel; currentLevel > level; currentLevel-- { - nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, currentLevel) + nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, currentLevel, visited) if err != nil { return err } @@ -373,7 +375,7 @@ func (i *HNSWIndex) insertBuiltNode(ctx context.Context, position int) error { } } for currentLevel := min(level, i.maxLevel); currentLevel >= 0; currentLevel-- { - candidates, err := i.searchHNSWLayer(ctx, query, []int{entry}, i.options.EFConstruction, currentLevel) + candidates, err := i.searchHNSWLayer(ctx, query, []int{entry}, i.options.EFConstruction, currentLevel, visited) if err != nil { return err } @@ -403,7 +405,7 @@ type hnswScoredNode struct { score float32 } -func (i *HNSWIndex) searchHNSWLayer(ctx context.Context, query []float32, entries []int, ef, level int) ([]hnswScoredNode, error) { +func (i *HNSWIndex) searchHNSWLayer(ctx context.Context, query []float32, entries []int, ef, level int, visited *hnswVisited) ([]hnswScoredNode, error) { limit := min(ef, len(i.keys)) if limit <= 0 { return []hnswScoredNode{}, nil @@ -412,9 +414,9 @@ func (i *HNSWIndex) searchHNSWLayer(ctx context.Context, query []float32, entrie worse := func(left, right hnswScoredNode) bool { return hnswNodeBetter(i.options.Metric, right, left) } candidates := ailego.NewHeap(better) results := ailego.NewHeap(worse) - visited := make([]bool, len(i.keys)) + visited.reset(len(i.keys)) for _, entry := range entries { - if entry < 0 || entry >= len(i.keys) || i.levels[entry] < level || visited[entry] { + if entry < 0 || entry >= len(i.keys) || i.levels[entry] < level || visited.seen(entry) { continue } score, err := i.computeDistance(query, i.vectorAt(entry)) @@ -422,7 +424,7 @@ func (i *HNSWIndex) searchHNSWLayer(ctx context.Context, query []float32, entrie return nil, err } node := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) candidates.Push(node) results.Push(node) } @@ -436,10 +438,10 @@ func (i *HNSWIndex) searchHNSWLayer(ctx context.Context, query []float32, entrie break } for _, neighbor := range i.neighbors[current.position][level] { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := i.computeDistance(query, i.vectorAt(neighbor)) if err != nil { return nil, err @@ -674,8 +676,10 @@ func (i *HNSWIndex) SearchHNSWGroups( return nil, err } entry := i.entryPoint + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) for level := i.maxLevel; level > 0; level-- { - nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, level) + nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, level, visited) if err != nil { return nil, fmt.Errorf("core: descend HNSW group-by level %d: %w", level, err) } @@ -689,7 +693,7 @@ func (i *HNSWIndex) SearchHNSWGroups( }, EF: options.EF, PrefetchOffset: options.PrefetchOffset, PrefetchLines: options.PrefetchLines, } - initial, err := i.searchHNSWBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions) + initial, err := i.searchHNSWBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions, visited) if err != nil { return nil, err } @@ -704,7 +708,7 @@ func (i *HNSWIndex) SearchHNSWGroups( } return expandHNSWGroups( ctx, i.options.Metric, i.keys, i.neighbors, initial, options.GroupByOptions, - scoreAt, func(score float32) float32 { return score }, groupNodeBetter(i.options.Metric, i.keys), prefetch, + scoreAt, func(score float32) float32 { return score }, groupNodeBetter(i.options.Metric, i.keys), prefetch, visited, ) } @@ -751,8 +755,10 @@ func (i *HNSWIndex) searchHNSW(ctx context.Context, query []float32, options HNS } entry := i.entryPoint + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) for level := i.maxLevel; level > 0; level-- { - nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, level) + nearest, err := i.searchHNSWLayer(ctx, query, []int{entry}, 1, level, visited) if err != nil { return nil, fmt.Errorf("core: descend HNSW level %d: %w", level, err) } @@ -761,7 +767,7 @@ func (i *HNSWIndex) searchHNSW(ctx context.Context, query []float32, options HNS } } capacity := max(options.EF, options.TopK) - candidates, err := i.searchHNSWBase(ctx, query, entry, capacity, options) + candidates, err := i.searchHNSWBase(ctx, query, entry, capacity, options, visited) if err != nil { return nil, err } @@ -775,19 +781,19 @@ func (i *HNSWIndex) searchHNSW(ctx context.Context, query []float32, options HNS return results, nil } -func (i *HNSWIndex) searchHNSWBase(ctx context.Context, query []float32, entry, capacity int, options HNSWSearchOptions) ([]hnswScoredNode, error) { +func (i *HNSWIndex) searchHNSWBase(ctx context.Context, query []float32, entry, capacity int, options HNSWSearchOptions, visited *hnswVisited) ([]hnswScoredNode, error) { better := func(left, right hnswScoredNode) bool { return hnswNodeBetter(i.options.Metric, left, right) } worse := func(left, right hnswScoredNode) bool { return i.hnswResultNodeBetter(right, left) } frontier := ailego.NewHeap(better) accepted := ailego.NewHeap(worse) - visited := make([]bool, len(i.keys)) + visited.reset(len(i.keys)) score, err := i.computeDistance(query, i.vectorAt(entry)) if err != nil { return nil, fmt.Errorf("core: score HNSW entry point: %w", err) } start := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) frontier.Push(start) if i.acceptHNSWResult(start, options.SearchOptions) { accepted.Push(start) @@ -805,10 +811,10 @@ func (i *HNSWIndex) searchHNSWBase(ctx context.Context, query []float32, entry, neighbors := i.neighbors[current.position][0] prefetchDenseHNSWNeighbors(i.vectors, i.dimension, neighbors, options.PrefetchOffset, options.PrefetchLines) for _, neighbor := range neighbors { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := i.computeDistance(query, i.vectorAt(neighbor)) if err != nil { return nil, fmt.Errorf("core: score HNSW node %d: %w", neighbor, err) diff --git a/internal/core/hnsw_parallel.go b/internal/core/hnsw_parallel.go index a784ab7..33138d0 100644 --- a/internal/core/hnsw_parallel.go +++ b/internal/core/hnsw_parallel.go @@ -66,10 +66,12 @@ func buildParallelHNSW( } func (g *parallelHNSWGraph) insert(ctx context.Context, position int) error { + visited := acquireHNSWVisited(len(g.levels)) + defer releaseHNSWVisited(visited) level := g.levels[position] entry, maxLevel := g.entrySnapshot() for currentLevel := maxLevel; currentLevel > level; currentLevel-- { - nearest, err := g.searchLayer(ctx, position, []int{entry}, 1, currentLevel) + nearest, err := g.searchLayer(ctx, position, []int{entry}, 1, currentLevel, visited) if err != nil { return err } @@ -78,7 +80,7 @@ func (g *parallelHNSWGraph) insert(ctx context.Context, position int) error { } } for currentLevel := min(level, maxLevel); currentLevel >= 0; currentLevel-- { - candidates, err := g.searchLayer(ctx, position, []int{entry}, g.options.EFConstruction, currentLevel) + candidates, err := g.searchLayer(ctx, position, []int{entry}, g.options.EFConstruction, currentLevel, visited) if err != nil { return err } @@ -119,6 +121,7 @@ func (g *parallelHNSWGraph) searchLayer( entries []int, ef int, level int, + visited *hnswVisited, ) ([]hnswScoredNode, error) { limit := min(ef, len(g.levels)) if limit <= 0 { @@ -128,9 +131,9 @@ func (g *parallelHNSWGraph) searchLayer( worse := func(left, right hnswScoredNode) bool { return hnswNodeBetter(g.options.Metric, right, left) } candidates := ailego.NewHeap(better) results := ailego.NewHeap(worse) - visited := make([]bool, len(g.levels)) + visited.reset(len(g.levels)) for _, entry := range entries { - if entry < 0 || entry >= len(g.levels) || g.levels[entry] < level || visited[entry] { + if entry < 0 || entry >= len(g.levels) || g.levels[entry] < level || visited.seen(entry) { continue } score, err := g.score(query, entry) @@ -138,7 +141,7 @@ func (g *parallelHNSWGraph) searchLayer( return nil, err } node := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) candidates.Push(node) results.Push(node) } @@ -153,10 +156,10 @@ func (g *parallelHNSWGraph) searchLayer( } g.nodeLocks[current.position].RLock() for _, neighbor := range g.neighbors[current.position][level] { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := g.score(query, neighbor) if err != nil { g.nodeLocks[current.position].RUnlock() diff --git a/internal/core/hnsw_quantized.go b/internal/core/hnsw_quantized.go index 237139d..dc33f41 100644 --- a/internal/core/hnsw_quantized.go +++ b/internal/core/hnsw_quantized.go @@ -171,9 +171,11 @@ func (i *ScalarQuantizedHNSWIndex) SearchHNSWGroups( scoreAt := func(position int) (float32, error) { return QuantizedDistance(i.vectors.metric, i.vectors.codes[position], queryCode) } + visited := acquireHNSWVisited(len(i.vectors.keys)) + defer releaseHNSWVisited(visited) entry := i.base.entryPoint for level := i.base.maxLevel; level > 0; level-- { - nearest, err := i.searchLayer(ctx, []int{entry}, 1, level, scoreAt) + nearest, err := i.searchLayer(ctx, []int{entry}, 1, level, scoreAt, visited) if err != nil { return nil, fmt.Errorf("core: descend scalar-quantized HNSW group-by level %d: %w", level, err) } @@ -187,7 +189,7 @@ func (i *ScalarQuantizedHNSWIndex) SearchHNSWGroups( }, EF: options.EF, PrefetchOffset: options.PrefetchOffset, PrefetchLines: options.PrefetchLines, } - initial, err := i.searchBase(ctx, entry, max(options.EF, candidateCount), searchOptions, scoreAt) + initial, err := i.searchBase(ctx, entry, max(options.EF, candidateCount), searchOptions, scoreAt, visited) if err != nil { return nil, err } @@ -199,7 +201,7 @@ func (i *ScalarQuantizedHNSWIndex) SearchHNSWGroups( } return expandHNSWGroups( ctx, i.vectors.metric, i.vectors.keys, i.base.neighbors, initial, options.GroupByOptions, - scoreAt, func(score float32) float32 { return score }, groupNodeBetter(i.vectors.metric, i.vectors.keys), prefetch, + scoreAt, func(score float32) float32 { return score }, groupNodeBetter(i.vectors.metric, i.vectors.keys), prefetch, visited, ) } @@ -251,9 +253,11 @@ func (i *ScalarQuantizedHNSWIndex) search( scoreAt := func(position int) (float32, error) { return QuantizedDistance(i.vectors.metric, i.vectors.codes[position], queryCode) } + visited := acquireHNSWVisited(len(i.vectors.keys)) + defer releaseHNSWVisited(visited) entry := i.base.entryPoint for level := i.base.maxLevel; level > 0; level-- { - nearest, err := i.searchLayer(ctx, []int{entry}, 1, level, scoreAt) + nearest, err := i.searchLayer(ctx, []int{entry}, 1, level, scoreAt, visited) if err != nil { return nil, fmt.Errorf("core: descend scalar-quantized HNSW level %d: %w", level, err) } @@ -262,7 +266,7 @@ func (i *ScalarQuantizedHNSWIndex) search( } } capacity := max(options.EF, options.TopK) - candidates, err := i.searchBase(ctx, entry, capacity, options, scoreAt) + candidates, err := i.searchBase(ctx, entry, capacity, options, scoreAt, visited) if err != nil { return nil, err } @@ -281,6 +285,7 @@ func (i *ScalarQuantizedHNSWIndex) searchLayer( entries []int, ef, level int, scoreAt func(int) (float32, error), + visited *hnswVisited, ) ([]hnswScoredNode, error) { limit := min(ef, len(i.vectors.keys)) if limit <= 0 { @@ -291,9 +296,9 @@ func (i *ScalarQuantizedHNSWIndex) searchLayer( worse := func(left, right hnswScoredNode) bool { return hnswNodeBetter(metric, right, left) } candidates := ailego.NewHeap(better) results := ailego.NewHeap(worse) - visited := make([]bool, len(i.vectors.keys)) + visited.reset(len(i.vectors.keys)) for _, entry := range entries { - if entry < 0 || entry >= len(i.vectors.keys) || i.base.levels[entry] < level || visited[entry] { + if entry < 0 || entry >= len(i.vectors.keys) || i.base.levels[entry] < level || visited.seen(entry) { continue } score, err := scoreAt(entry) @@ -301,7 +306,7 @@ func (i *ScalarQuantizedHNSWIndex) searchLayer( return nil, err } node := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) candidates.Push(node) results.Push(node) } @@ -315,10 +320,10 @@ func (i *ScalarQuantizedHNSWIndex) searchLayer( break } for _, neighbor := range i.base.neighbors[current.position][level] { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := scoreAt(neighbor) if err != nil { return nil, err @@ -352,20 +357,21 @@ func (i *ScalarQuantizedHNSWIndex) searchBase( entry, capacity int, options HNSWSearchOptions, scoreAt func(int) (float32, error), + visited *hnswVisited, ) ([]hnswScoredNode, error) { metric := i.vectors.metric better := func(left, right hnswScoredNode) bool { return hnswNodeBetter(metric, left, right) } worse := func(left, right hnswScoredNode) bool { return i.resultNodeBetter(right, left) } frontier := ailego.NewHeap(better) accepted := ailego.NewHeap(worse) - visited := make([]bool, len(i.vectors.keys)) + visited.reset(len(i.vectors.keys)) score, err := scoreAt(entry) if err != nil { return nil, fmt.Errorf("core: score scalar-quantized HNSW entry point: %w", err) } start := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) frontier.Push(start) if i.acceptResult(start, options.SearchOptions) { accepted.Push(start) @@ -382,10 +388,10 @@ func (i *ScalarQuantizedHNSWIndex) searchBase( neighbors := i.base.neighbors[current.position][0] prefetchQuantizedHNSWNeighbors(i.vectors.codes, neighbors, options.PrefetchOffset, options.PrefetchLines) for _, neighbor := range neighbors { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := scoreAt(neighbor) if err != nil { return nil, fmt.Errorf("core: score scalar-quantized HNSW node %d: %w", neighbor, err) diff --git a/internal/core/hnsw_rabitq.go b/internal/core/hnsw_rabitq.go index 17791f2..61ba7ff 100644 --- a/internal/core/hnsw_rabitq.go +++ b/internal/core/hnsw_rabitq.go @@ -440,6 +440,8 @@ func (i *HNSWRaBitQIndex) SearchHNSWRaBitQGroups( if err != nil { return nil, err } + visited := acquireHNSWVisited(len(i.codes)) + defer releaseHNSWVisited(visited) entry := i.base.entryPoint for level := i.base.maxLevel; level > 0; level-- { entry, err = i.searchRaBitQLayer(ctx, query, entry, level) @@ -450,7 +452,7 @@ func (i *HNSWRaBitQIndex) SearchHNSWRaBitQGroups( searchOptions := SearchOptions{ TopK: candidateCount, Radius: options.Radius, Filter: options.Filter, } - initial, err := i.searchRaBitQBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions) + initial, err := i.searchRaBitQBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions, visited) if err != nil { return nil, err } @@ -472,7 +474,7 @@ func (i *HNSWRaBitQIndex) SearchHNSWRaBitQGroups( } return expandHNSWGroups( ctx, i.options.Metric, i.base.keys, i.base.neighbors, initial, options.GroupByOptions, - scoreAt, i.publicRaBitQScore, better, nil, + scoreAt, i.publicRaBitQScore, better, nil, visited, ) } @@ -588,6 +590,8 @@ func (i *HNSWRaBitQIndex) searchPrepared(ctx context.Context, query *RaBitQQuery if linear || len(i.codes) <= DefaultHNSWBruteForceThreshold { return i.scanRaBitQCodes(ctx, query, options) } + visited := acquireHNSWVisited(len(i.codes)) + defer releaseHNSWVisited(visited) entry := i.base.entryPoint for level := i.base.maxLevel; level > 0; level-- { nearest, err := i.searchRaBitQLayer(ctx, query, entry, level) @@ -596,7 +600,7 @@ func (i *HNSWRaBitQIndex) searchPrepared(ctx context.Context, query *RaBitQQuery } entry = nearest } - nodes, err := i.searchRaBitQBase(ctx, query, entry, max(ef, options.TopK), options) + nodes, err := i.searchRaBitQBase(ctx, query, entry, max(ef, options.TopK), options, visited) if err != nil { return nil, err } @@ -662,18 +666,18 @@ func (i *HNSWRaBitQIndex) searchRaBitQLayer(ctx context.Context, query *RaBitQQu } } -func (i *HNSWRaBitQIndex) searchRaBitQBase(ctx context.Context, query *RaBitQQuery, entry, capacity int, options SearchOptions) ([]hnswScoredNode, error) { +func (i *HNSWRaBitQIndex) searchRaBitQBase(ctx context.Context, query *RaBitQQuery, entry, capacity int, options SearchOptions, visited *hnswVisited) ([]hnswScoredNode, error) { better := func(left, right hnswScoredNode) bool { return i.raBitQNodeBetter(left, right) } worse := func(left, right hnswScoredNode) bool { return i.raBitQNodeBetter(right, left) } frontier := ailego.NewHeap(better) accepted := ailego.NewHeap(worse) - visited := make([]bool, len(i.codes)) + visited.reset(len(i.codes)) estimate, err := query.Estimate(i.codes[entry]) if err != nil { return nil, err } start := hnswScoredNode{position: entry, score: estimate.Distance} - visited[entry] = true + visited.mark(entry) frontier.Push(start) if i.acceptRaBitQNode(start, options) { accepted.Push(start) @@ -688,10 +692,10 @@ func (i *HNSWRaBitQIndex) searchRaBitQBase(ctx context.Context, query *RaBitQQue break } for _, neighbor := range i.base.neighbors[current.position][0] { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) coarse, err := query.EstimateCoarse(i.codes[neighbor]) if err != nil { return nil, err diff --git a/internal/core/hnsw_sparse.go b/internal/core/hnsw_sparse.go index be82b0d..4d12064 100644 --- a/internal/core/hnsw_sparse.go +++ b/internal/core/hnsw_sparse.go @@ -313,6 +313,8 @@ func (i *SparseHNSWIndex) Neighbors(key uint64, level int) ([]uint64, error) { } func (i *SparseHNSWIndex) insertBuiltNode(ctx context.Context, position int) error { + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) level := i.levels[position] if i.entryPoint < 0 { i.entryPoint = position @@ -322,7 +324,7 @@ func (i *SparseHNSWIndex) insertBuiltNode(ctx context.Context, position int) err entry := i.entryPoint query := i.sparseVectorAt(position) for currentLevel := i.maxLevel; currentLevel > level; currentLevel-- { - nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, currentLevel) + nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, currentLevel, visited) if err != nil { return err } @@ -331,7 +333,7 @@ func (i *SparseHNSWIndex) insertBuiltNode(ctx context.Context, position int) err } } for currentLevel := min(level, i.maxLevel); currentLevel >= 0; currentLevel-- { - candidates, err := i.searchLayer(ctx, query, []int{entry}, i.options.EFConstruction, currentLevel) + candidates, err := i.searchLayer(ctx, query, []int{entry}, i.options.EFConstruction, currentLevel, visited) if err != nil { return err } @@ -356,7 +358,7 @@ func (i *SparseHNSWIndex) insertBuiltNode(ctx context.Context, position int) err return nil } -func (i *SparseHNSWIndex) searchLayer(ctx context.Context, query SparseVector, entries []int, ef, level int) ([]hnswScoredNode, error) { +func (i *SparseHNSWIndex) searchLayer(ctx context.Context, query SparseVector, entries []int, ef, level int, visited *hnswVisited) ([]hnswScoredNode, error) { limit := min(ef, len(i.keys)) if limit <= 0 { return []hnswScoredNode{}, nil @@ -365,9 +367,9 @@ func (i *SparseHNSWIndex) searchLayer(ctx context.Context, query SparseVector, e worse := func(left, right hnswScoredNode) bool { return hnswNodeBetter(MetricIP, right, left) } candidates := ailego.NewHeap(better) results := ailego.NewHeap(worse) - visited := make([]bool, len(i.keys)) + visited.reset(len(i.keys)) for _, entry := range entries { - if entry < 0 || entry >= len(i.keys) || i.levels[entry] < level || visited[entry] { + if entry < 0 || entry >= len(i.keys) || i.levels[entry] < level || visited.seen(entry) { continue } score, err := sparseHNSWScore(query, i.sparseVectorAt(entry)) @@ -375,7 +377,7 @@ func (i *SparseHNSWIndex) searchLayer(ctx context.Context, query SparseVector, e return nil, err } node := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) candidates.Push(node) results.Push(node) } @@ -389,10 +391,10 @@ func (i *SparseHNSWIndex) searchLayer(ctx context.Context, query SparseVector, e break } for _, neighbor := range i.neighbors[current.position][level] { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := sparseHNSWScore(query, i.sparseVectorAt(neighbor)) if err != nil { return nil, err @@ -570,8 +572,10 @@ func (i *SparseHNSWIndex) SearchSparseHNSWGroups( return nil, err } entry := i.entryPoint + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) for level := i.maxLevel; level > 0; level-- { - nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, level) + nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, level, visited) if err != nil { return nil, fmt.Errorf("core: descend sparse HNSW group-by level %d: %w", level, err) } @@ -585,7 +589,7 @@ func (i *SparseHNSWIndex) SearchSparseHNSWGroups( }, EF: options.EF, PrefetchOffset: options.PrefetchOffset, PrefetchLines: options.PrefetchLines, } - initial, err := i.searchBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions) + initial, err := i.searchBase(ctx, query, entry, max(options.EF, candidateCount), searchOptions, visited) if err != nil { return nil, err } @@ -600,7 +604,7 @@ func (i *SparseHNSWIndex) SearchSparseHNSWGroups( } return expandHNSWGroups( ctx, MetricIP, i.keys, i.neighbors, initial, options.GroupByOptions, - scoreAt, func(score float32) float32 { return score }, groupNodeBetter(MetricIP, i.keys), prefetch, + scoreAt, func(score float32) float32 { return score }, groupNodeBetter(MetricIP, i.keys), prefetch, visited, ) } @@ -642,8 +646,10 @@ func (i *SparseHNSWIndex) searchSparseHNSW(ctx context.Context, query SparseVect } entry := i.entryPoint + visited := acquireHNSWVisited(len(i.keys)) + defer releaseHNSWVisited(visited) for level := i.maxLevel; level > 0; level-- { - nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, level) + nearest, err := i.searchLayer(ctx, query, []int{entry}, 1, level, visited) if err != nil { return nil, fmt.Errorf("core: descend sparse HNSW level %d: %w", level, err) } @@ -652,7 +658,7 @@ func (i *SparseHNSWIndex) searchSparseHNSW(ctx context.Context, query SparseVect } } capacity := max(options.EF, options.TopK) - candidates, err := i.searchBase(ctx, query, entry, capacity, options) + candidates, err := i.searchBase(ctx, query, entry, capacity, options, visited) if err != nil { return nil, err } @@ -666,19 +672,19 @@ func (i *SparseHNSWIndex) searchSparseHNSW(ctx context.Context, query SparseVect return results, nil } -func (i *SparseHNSWIndex) searchBase(ctx context.Context, query SparseVector, entry, capacity int, options HNSWSearchOptions) ([]hnswScoredNode, error) { +func (i *SparseHNSWIndex) searchBase(ctx context.Context, query SparseVector, entry, capacity int, options HNSWSearchOptions, visited *hnswVisited) ([]hnswScoredNode, error) { better := func(left, right hnswScoredNode) bool { return hnswNodeBetter(MetricIP, left, right) } worse := func(left, right hnswScoredNode) bool { return i.resultNodeBetter(right, left) } frontier := ailego.NewHeap(better) accepted := ailego.NewHeap(worse) - visited := make([]bool, len(i.keys)) + visited.reset(len(i.keys)) score, err := sparseHNSWScore(query, i.sparseVectorAt(entry)) if err != nil { return nil, fmt.Errorf("core: score sparse HNSW entry point: %w", err) } start := hnswScoredNode{position: entry, score: score} - visited[entry] = true + visited.mark(entry) frontier.Push(start) if i.acceptResult(start, options.SearchOptions) { accepted.Push(start) @@ -696,10 +702,10 @@ func (i *SparseHNSWIndex) searchBase(ctx context.Context, query SparseVector, en neighbors := i.neighbors[current.position][0] prefetchSparseHNSWNeighbors(i.offsets, i.indices, i.values, neighbors, options.PrefetchOffset, options.PrefetchLines) for _, neighbor := range neighbors { - if visited[neighbor] { + if visited.seen(neighbor) { continue } - visited[neighbor] = true + visited.mark(neighbor) score, err := sparseHNSWScore(query, i.sparseVectorAt(neighbor)) if err != nil { return nil, fmt.Errorf("core: score sparse HNSW node %d: %w", neighbor, err) diff --git a/internal/core/hnsw_test.go b/internal/core/hnsw_test.go index 23324c8..c378cdd 100644 --- a/internal/core/hnsw_test.go +++ b/internal/core/hnsw_test.go @@ -770,6 +770,7 @@ func BenchmarkHNSWSearch(b *testing.B) { index := buildSearchHNSW(b, MetricL2, inputs, 16, 120) query := inputs[4321].Vector options := HNSWSearchOptions{SearchOptions: SearchOptions{TopK: 10}, EF: 100} + b.ReportAllocs() b.ResetTimer() for b.Loop() { { diff --git a/internal/core/hnsw_visited.go b/internal/core/hnsw_visited.go new file mode 100644 index 0000000..52940c6 --- /dev/null +++ b/internal/core/hnsw_visited.go @@ -0,0 +1,51 @@ +// SPDX-License-Identifier: Apache-2.0 + +package core + +import "sync" + +var hnswVisitedPool = sync.Pool{ + New: func() any { return new(hnswVisited) }, +} + +func acquireHNSWVisited(size int) *hnswVisited { + visited := hnswVisitedPool.Get().(*hnswVisited) + visited.reset(size) + return visited +} + +func releaseHNSWVisited(visited *hnswVisited) { + hnswVisitedPool.Put(visited) +} + +// hnswVisited tracks graph visits without clearing the full node-sized buffer +// between traversals. A generation value distinguishes marks from consecutive +// traversals; the buffer is cleared only when the byte generation wraps. +type hnswVisited struct { + marks []uint8 + generation uint8 +} + +func (v *hnswVisited) reset(size int) { + if cap(v.marks) < size { + v.marks = make([]uint8, size) + v.generation = 1 + return + } + v.marks = v.marks[:size] + v.generation++ + if v.generation == 0 { + v.marks = v.marks[:cap(v.marks)] + clear(v.marks) + v.marks = v.marks[:size] + v.generation = 1 + } +} + +func (v *hnswVisited) seen(position int) bool { + return v.marks[position] == v.generation +} + +func (v *hnswVisited) mark(position int) { + v.marks[position] = v.generation +} diff --git a/internal/core/hnsw_visited_test.go b/internal/core/hnsw_visited_test.go new file mode 100644 index 0000000..ce81450 --- /dev/null +++ b/internal/core/hnsw_visited_test.go @@ -0,0 +1,80 @@ +// SPDX-License-Identifier: Apache-2.0 + +package core + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestHNSWVisitedReset(t *testing.T) { + var visited hnswVisited + visited.reset(4) + visited.mark(1) + require.True(t, visited.seen(1)) + require.False(t, visited.seen(2)) + + visited.reset(4) + require.False(t, visited.seen(1)) + visited.mark(2) + require.True(t, visited.seen(2)) +} + +func TestHNSWVisitedResize(t *testing.T) { + var visited hnswVisited + visited.reset(2) + visited.mark(1) + visited.reset(5) + + require.Len(t, visited.marks, 5) + for position := range visited.marks { + require.False(t, visited.seen(position)) + } + visited.mark(4) + require.True(t, visited.seen(4)) +} + +func TestHNSWVisitedGenerationWrap(t *testing.T) { + var visited hnswVisited + visited.reset(3) + visited.mark(1) + + for range 255 { + visited.reset(3) + } + + require.NotZero(t, visited.generation) + for position := range visited.marks { + require.False(t, visited.seen(position)) + } +} + +func TestHNSWVisitedGenerationWrapAfterShrink(t *testing.T) { + var visited hnswVisited + visited.reset(4) + visited.mark(3) + visited.marks[3] = 2 // Simulate a stale mark matching the first generation after wrap. + + visited.marks = visited.marks[:2] + for range 255 { + visited.reset(2) + } + visited.reset(4) + + require.False(t, visited.seen(3)) +} + +func TestHNSWVisitedPoolReusesAllocation(t *testing.T) { + visited := acquireHNSWVisited(1024) + releaseHNSWVisited(visited) + allocations := testing.AllocsPerRun(100, func() { + visited := acquireHNSWVisited(1024) + visited.mark(100) + if !visited.seen(100) { + panic("marked position is not visited") + } + releaseHNSWVisited(visited) + }) + require.Zero(t, allocations) +}