search method

List<VectorSearchHit> search(
  1. Vector query,
  2. int k, {
  3. VectorMetric? metric,
  4. int? nprobe,
})

Top-k nearest neighbors of query under metric (or defaultMetric). Optional per-call nprobe overrides the field.

Implementation

List<VectorSearchHit> search(
  Vector query,
  int k, {
  VectorMetric? metric,
  int? nprobe,
}) {
  if (!_trained) {
    throw StateError('IvfFlatIndex.search: call train() first');
  }
  if (query.dim != dim) {
    throw StateError(
      'IvfFlatIndex.search: query dim ${query.dim} != index dim $dim',
    );
  }
  if (k <= 0 || _n == 0) return const [];
  final m = metric ?? defaultMetric;
  final larger = m == VectorMetric.innerProduct;
  final cells = _probeCells(query, nprobe ?? this.nprobe);

  // Precompute cosine query-norm.
  double qNorm = 0.0;
  if (m == VectorMetric.cosine) {
    for (final x in query.values) {
      qNorm += x * x;
    }
    qNorm = math.sqrt(qNorm);
    if (qNorm == 0.0) qNorm = 1.0;
  }

  final effK = math.min(k, _n);
  final scores = List<double>.filled(effK, 0.0);
  final ids = List<Object?>.filled(effK, null);
  var filled = 0;

  for (final cell in cells) {
    final vecs = _cellVecs[cell];
    final cellIds = _cellIds[cell];
    final count = _cellCounts[cell];
    for (var r = 0; r < count; r++) {
      final base = r * dim;
      double score;
      switch (m) {
        case VectorMetric.l2sq:
        case VectorMetric.l2:
          var s = 0.0;
          for (var i = 0; i < dim; i++) {
            final d = vecs[base + i] - query.values[i];
            s += d * d;
          }
          score = s;
          break;
        case VectorMetric.innerProduct:
          var s = 0.0;
          for (var i = 0; i < dim; i++) {
            s += vecs[base + i] * query.values[i];
          }
          score = s;
          break;
        case VectorMetric.cosine:
          var dot = 0.0, norm = 0.0;
          for (var i = 0; i < dim; i++) {
            final a = vecs[base + i];
            dot += a * query.values[i];
            norm += a * a;
          }
          norm = math.sqrt(norm);
          score = norm == 0.0 ? 1.0 : 1.0 - dot / (norm * qNorm);
          break;
      }
      _insertTopK(scores, ids, filled, score, cellIds[r], larger, effK);
      if (filled < effK) filled++;
    }
  }

  final out = <VectorSearchHit>[];
  for (var i = 0; i < filled; i++) {
    final s = m == VectorMetric.l2 ? math.sqrt(scores[i]) : scores[i];
    out.add(VectorSearchHit(ids[i], s));
  }
  return out;
}