Skip to content
Open
Show file tree
Hide file tree
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
15 changes: 10 additions & 5 deletions graph.go
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,10 @@ func (g *Graph[K]) Add(nodes ...Node[K]) {

g.assertDims(vec)
insertLevel := g.randomLevel()
if _, ok := g.Lookup(key); ok {
g.Delete(key)
}

// Create layers that don't exist yet.
for insertLevel >= len(g.layers) {
g.layers = append(g.layers, &layer[K]{})
Expand Down Expand Up @@ -388,7 +392,9 @@ func (g *Graph[K]) Add(nodes ...Node[K]) {
// On subsequent layers, we use the elevator node to enter the graph
// at the best point.
if elevator != nil {
searchPoint = layer.nodes[*elevator]
if next, ok := layer.nodes[*elevator]; ok {
searchPoint = next
}
}

if g.Distance == nil {
Expand All @@ -406,9 +412,6 @@ func (g *Graph[K]) Add(nodes ...Node[K]) {
elevator = ptr(neighborhood[0].node.Key)

if insertLevel >= i {
if _, ok := layer.nodes[key]; ok {
g.Delete(key)
}
// Insert the new node into the layer.
layer.nodes[key] = newNode
for _, node := range neighborhood {
Expand Down Expand Up @@ -464,7 +467,9 @@ func (h *Graph[K]) search(near Vector, k int) []SearchResult[K] {
for layer := len(h.layers) - 1; layer >= 0; layer-- {
searchPoint := h.layers[layer].entry()
if elevator != nil {
searchPoint = h.layers[layer].nodes[*elevator]
if next, ok := h.layers[layer].nodes[*elevator]; ok {
searchPoint = next
}
}

// Descending hierarchies
Expand Down
38 changes: 38 additions & 0 deletions graph_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,44 @@ func TestGraph_AddDelete(t *testing.T) {
})
}

func TestGraph_AddReplacesExistingNode(t *testing.T) {
t.Parallel()

const numNodes = 128
g := newTestGraph[int]()
for i := 0; i < numNodes; i++ {
g.Add(MakeNode(i, Vector{float32(i)}))
}

for i := 0; i < numNodes; i++ {
replacement := Vector{float32(i)}
g.Add(MakeNode(i, replacement))

require.Equal(t, numNodes, g.Len())
got, ok := g.Lookup(i)
require.True(t, ok)
require.Equal(t, replacement, got)
}

for i := 0; i < numNodes; i++ {
replacement := Vector{float32(numNodes + i)}
g.Add(MakeNode(i, replacement))

require.Equal(t, numNodes, g.Len())
got, ok := g.Lookup(i)
require.True(t, ok)
require.Equal(t, replacement, got)
}

for i := 0; i < numNodes; i++ {
replacement := Vector{float32(numNodes + i)}
results := g.Search(replacement, 1)
require.NotEmpty(t, results)
require.Equal(t, i, results[0].Key)
require.Equal(t, replacement, results[0].Value)
}
}

func Benchmark_HSNW(b *testing.B) {
b.ReportAllocs()

Expand Down
Loading