diff --git a/sql/vector_mhnsw.cc b/sql/vector_mhnsw.cc index c480c36c7e7ad..0805e1d7104a7 100644 --- a/sql/vector_mhnsw.cc +++ b/sql/vector_mhnsw.cc @@ -1125,11 +1125,11 @@ class VisitedSet { auto *v= new (root) Visited(node, dist); insert(node); - count++; return v; } void insert(const FVectorNode *n) { + count++; nodes[idx++]= n; if (idx == 8) flush(); } @@ -1355,10 +1355,14 @@ static int search_layer(MHNSW_param *p, const FVector *target, float threshold, return err; if (!best.is_full()) { - Visited *v= visited.create(links[i], links[i]->distance_to(target)); - if (v->distance_to_target <= threshold) + float distance= links[i]->distance_to(target); + if (distance <= threshold) + { + visited.insert(links[i]); continue; - p->acc.diameter= std::max(p->acc.diameter, v->distance_to_target); + } + p->acc.diameter= std::max(p->acc.diameter, distance); + Visited *v= visited.create(links[i], distance); candidates.safe_push(v); if (skip_deleted && v->node->deleted) continue; @@ -1367,21 +1371,21 @@ static int search_layer(MHNSW_param *p, const FVector *target, float threshold, } else { - Visited *v= visited.create(links[i], - links[i]->distance_greater_than(target, furthest_best, - p->mode, &p->acc)); - if (v->distance_to_target <= threshold) + float distance= links[i]->distance_greater_than(target, furthest_best, + p->mode, &p->acc); + if (distance <= threshold || distance >= furthest_best) + { + visited.insert(links[i]); + continue; + } + Visited *v= visited.create(links[i], distance); + candidates.safe_push(v); + if (skip_deleted && v->node->deleted) continue; - if (v->distance_to_target < furthest_best) + if (distance < best.top()->distance_to_target) { - candidates.safe_push(v); - if (skip_deleted && v->node->deleted) - continue; - if (v->distance_to_target < best.top()->distance_to_target) - { - best.replace_top(v); - furthest_best= lenient_furthest(best, p->acc.diameter, leniency); - } + best.replace_top(v); + furthest_best= lenient_furthest(best, p->acc.diameter, leniency); } } }