gsmiller commented on code in PR #16709: URL: https://github.com/apache/lucene/pull/16709#discussion_r4148452826
########## lucene/sandbox/src/java/org/apache/lucene/sandbox/codecs/segmentivf/Centroids.java: ########## @@ -0,0 +1,375 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.lucene.sandbox.codecs.segmentivf; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Random; +import org.apache.lucene.sandbox.codecs.segmentivf.Clustering.Parallel; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.CodeRecord; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.FineCodec; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.Nitrox2; +import org.apache.lucene.store.ByteArrayDataInput; +import org.apache.lucene.store.IndexOutput; +import org.apache.lucene.store.RandomAccessInput; +import org.apache.lucene.util.ArrayUtil; +import org.apache.lucene.util.VectorUtil; + +/** Encodes, ranks, and navigates the centroids used to route vectors into IVF cells. */ +final class Centroids { + private Centroids() {} + + /** + * Keeps centroid representations together so routing can shortlist cheaply and verify accurately. + */ + static final class CentroidCodes { + private static final int TILE = 512; + + final float[][] centroids; + final byte[] coarse; + final int nlist, coarseBytes; + private final int dim, fineStride; + private final FineCodec fine; + private final byte[] fineRecords; + private final ThreadLocal<byte[]> gathered = ThreadLocal.withInitial(() -> new byte[0]); + + CentroidCodes(float[][] centroids, int dim, FineCodec fine) { + this.dim = dim; + this.nlist = centroids.length; + this.centroids = centroids; + this.coarseBytes = Nitrox2.bytesPerVector(dim); + this.coarse = new byte[nlist * coarseBytes]; + this.fine = fine; + this.fineStride = fine == null ? 0 : CodeRecord.length(fine.codeBytes); + this.fineRecords = fine == null ? null : new byte[nlist * fineStride]; + encodeAll(); + } + + void encodeAll() { + for (int c = 0; c < nlist; c++) { + Nitrox2.encode(centroids[c], dim, coarse, c * coarseBytes); + if (fine != null) fine.encode(centroids[c], fineRecords, c * fineStride); + } + } + + void rankCandidates(float[] vector, int[] cands, int count, float[] out) { + if (fine == null) { + for (int i = 0; i < count; i++) out[i] = exactDistance(vector, cands[i]); + return; + } + byte[] flat = gathered.get(); + if (flat.length < count * fineStride) gathered.set(flat = new byte[count * fineStride]); + for (int i = 0; i < count; i++) { + System.arraycopy(fineRecords, cands[i] * fineStride, flat, i * fineStride, fineStride); + } + fine.query(vector, null).score(flat, fineStride, count, out); + for (int i = 0; i < count; i++) out[i] = -out[i]; + } + + static final class Routing { + final int[] cells; + int count, cell2; + float d1, d2; + + Routing(int capacity) { + cells = new int[capacity]; + } + } + + static final class Scratch { + final int[] coarseDist = new int[TILE], verifyCells; + final long[] heap; + final float[] verifyDist; + final byte[] qCode; + + Scratch(int dim, int nlist, int shortlist) { + heap = new long[shortlist]; + verifyCells = new int[shortlist]; + verifyDist = new float[shortlist]; + qCode = new byte[Nitrox2.bytesPerVector(dim)]; + } + } + + void routePacked(float[] vector, int shortlist, int keep, Routing out, Scratch scratch) { + final int want = Math.min(shortlist, nlist); + final long[] heap = scratch.heap; + final int[] cd = scratch.coarseDist; + int n = 0, worst = Integer.MAX_VALUE; + for (int base = 0; base < nlist; base += TILE) { + final int rows = Math.min(TILE, nlist - base); + Kernels.INSTANCE.hamming(scratch.qCode, coarse, base * coarseBytes, rows, cd); + for (int r = 0; r < rows; r++) { + final int dist = cd[r]; + if (n < want) { + heap[n++] = ((long) dist << 32) | (base + r); + if (n == want) { + for (int h = (n >>> 1) - 1; h >= 0; h--) { + CentroidGraph.siftDown(heap, h, n, heap[h], true); + } + worst = (int) (heap[0] >>> 32); + } + } else if (dist < worst) { + CentroidGraph.siftDown(heap, 0, n, ((long) dist << 32) | (base + r), true); + worst = (int) (heap[0] >>> 32); + } + } + } + final int[] cells = scratch.verifyCells; + final float[] dists = scratch.verifyDist; + for (int i = 0; i < n; i++) { + cells[i] = (int) heap[i]; + dists[i] = exactDistance(vector, cells[i]); + } + CentroidGraph.sortByDistance(dists, cells, n); + System.arraycopy(cells, 0, out.cells, 0, out.count = Math.min(keep, n)); + out.cell2 = n > 1 ? cells[1] : -1; + out.d1 = n > 0 ? dists[0] : Float.MAX_VALUE; + out.d2 = n > 1 ? dists[1] : Float.MAX_VALUE; + } + + float exactDistance(float[] vector, int c) { + return -VectorUtil.dotProduct(vector, centroids[c]); + } + + static boolean withinMargin(float d1, float d2, float margin) { + return d2 != Float.MAX_VALUE && d2 - d1 <= (margin - 1f) * Math.abs(d1); + } + } + + /** + * A compact graph over centroid codes that avoids scoring every cell when selecting query probes. + */ + record CentroidGraph( + int nlist, int coarseBytes, int stride, int entry, byte[] nodes, int[][] building) { + static final int M = 16, EF_CONSTRUCTION = 64, EF_MULTIPLIER = 2, MIN_EF = 32; + private static final int ALIGN = 64, ORD_BYTES = 2, LOCK_STRIPES = 512, INSERT_GRAIN = 256; + private static final ThreadLocal<int[]> VISITED = ThreadLocal.withInitial(() -> new int[1]); + + static CentroidGraph build(CentroidCodes codes, int dim) throws IOException { + int nlist = codes.nlist, coarseBytes = codes.coarseBytes; + int stride = (coarseBytes + 2 + M * ORD_BYTES + ALIGN - 1) / ALIGN * ALIGN; + int[][] neighbours = new int[nlist][]; + Arrays.fill(neighbours, new int[0]); + int[] order = new int[nlist]; + for (int i = 0; i < nlist; i++) order[i] = i; + Random random = new Random(0x5DEECE66DL); + for (int i = nlist - 1; i > 0; i--) { + int j = random.nextInt(i + 1), t = order[i]; + order[i] = order[j]; + order[j] = t; + } + int entry = order[0]; + CentroidGraph partial = + new CentroidGraph(nlist, coarseBytes, coarseBytes, entry, codes.coarse, neighbours); + Object[] locks = new Object[LOCK_STRIPES]; + for (int i = 0; i < LOCK_STRIPES; i++) locks[i] = new Object(); + Parallel.RangeTask insertRange = + (from, to) -> { + byte[] code = new byte[coarseBytes]; + int[] visited = new int[nlist]; + for (int idx = from + 1; idx <= to; idx++) { + int node = order[idx]; + Nitrox2.encode(codes.centroids[node], dim, code, 0); + int n = partial.search(code, EF_CONSTRUCTION, visited, null); + int[] kept = neighbours[node] = prune(codes, codes.centroids[node], visited, n); + for (int i = 0; i < kept.length; i++) { + int x = kept[i]; + synchronized (locks[(x * 0x9E3779B9) >>> 1 & (LOCK_STRIPES - 1)]) { + neighbours[x] = link(codes, neighbours[x], x, node, i == 0); + } + } + } + }; + int seed = Math.min(nlist - 1, Math.max(64, M * 4)); + insertRange.run(0, seed); + Parallel.overRange( + nlist - 1 - seed, INSERT_GRAIN, (from, to) -> insertRange.run(seed + from, seed + to)); + connect(neighbours, entry); + byte[] nodes = new byte[nlist * stride]; + for (int c = 0; c < nlist; c++) { + System.arraycopy(codes.coarse, c * coarseBytes, nodes, c * stride, coarseBytes); + int off = c * stride + coarseBytes, deg = Math.min(M, neighbours[c].length); + nodes[off] = (byte) deg; + for (int i = 0; i < deg; i++) { + nodes[off += ORD_BYTES] = (byte) neighbours[c][i]; + nodes[off + 1] = (byte) (neighbours[c][i] >>> 8); + } + } + return new CentroidGraph(nlist, coarseBytes, stride, entry, nodes, null); + } + + private static void connect(int[][] neighbours, int entry) { + int n = neighbours.length; + boolean[] reachable = new boolean[n]; + int[] queue = new int[n]; + for (int guard = 0; guard <= 8; guard++) { + Arrays.fill(reachable, false); + int head = 0, tail = 0; + queue[tail++] = entry; + reachable[entry] = true; + while (head < tail) { + for (int x : neighbours[queue[head++]]) { + if (reachable[x]) continue; + reachable[x] = true; + queue[tail++] = x; + } + } + if (tail == n || guard == 8) return; + for (int c = 0, host = entry; c < n; host = c++) { + if (reachable[c]) continue; + neighbours[c] = appendUnique(neighbours[c], host); + if (neighbours[host].length < M) { + neighbours[host] = appendUnique(neighbours[host], c); + } else { + neighbours[host] = neighbours[host].clone(); + neighbours[host][M - 1] = c; + } + reachable[c] = true; + } + } + } + + private static int[] appendUnique(int[] a, int v) { + for (int x : a) if (x == v) return a; + if (a.length >= M) return a; + int[] out = ArrayUtil.growExact(a, a.length + 1); + out[a.length] = v; + return out; + } + + private static void sortByDistance(float[] dist, int[] ids, int n) { + for (int i = 1; i < n; i++) { + float d = dist[i]; + int c = ids[i], j = i - 1; + for (; j >= 0 && dist[j] > d; j--) { + dist[j + 1] = dist[j]; + ids[j + 1] = ids[j]; + } + dist[j + 1] = d; + ids[j + 1] = c; + } + } + + private static int[] prune(CentroidCodes codes, float[] vec, int[] cand, int n) { + float[] dist = new float[n]; + for (int i = 0; i < n; i++) dist[i] = codes.exactDistance(vec, cand[i]); + sortByDistance(dist, cand, n); + int[] kept = new int[Math.min(M, n)]; + int nKept = 0; + for (int i = 0; i < n && nKept < kept.length; i++) { + boolean diverse = true; + for (int k = 0; k < nKept && diverse; k++) { + diverse = (codes.exactDistance(codes.centroids[cand[i]], kept[k]) < dist[i]) == false; + } + if (diverse) kept[nKept++] = cand[i]; + } + return ArrayUtil.copyOfSubArray(kept, 0, nKept); + } + + private static int[] link(CentroidCodes codes, int[] cur, int x, int node, boolean mustLink) { + if (cur.length < M) return appendUnique(cur, node); + for (int y : cur) if (y == node) return cur; + float[] xVec = codes.centroids[x]; + int worst = -1; + float worstD = Float.NEGATIVE_INFINITY; + for (int i = 0; i < cur.length; i++) { + float d = codes.exactDistance(xVec, cur[i]); + if (d > worstD) { + worstD = d; + worst = i; + } + } + if ((mustLink || codes.exactDistance(xVec, node) < worstD) == false) return cur; + int[] out = cur.clone(); + out[worst] = node; + return out; + } + + int search(byte[] qCode, int ef, int[] out, int[] outDist) { + int[] visited = VISITED.get(); + if (visited.length <= nlist) VISITED.set(visited = new int[nlist + 1]); + int gen = ++visited[0], nOut = 0, frontierN = 0, bestN = 0; + int cap = Math.min(Math.max(ef, 1), nlist); + long[] frontier = new long[Math.min(nlist, Math.max(64, cap * 4))], best = new long[cap]; Review Comment: I think you could use `LongHeap` off-the-shelf for `frontier` and `best`? ########## lucene/sandbox/src/java/org/apache/lucene/sandbox/codecs/segmentivf/Centroids.java: ########## @@ -0,0 +1,375 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.lucene.sandbox.codecs.segmentivf; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Random; +import org.apache.lucene.sandbox.codecs.segmentivf.Clustering.Parallel; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.CodeRecord; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.FineCodec; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.Nitrox2; +import org.apache.lucene.store.ByteArrayDataInput; +import org.apache.lucene.store.IndexOutput; +import org.apache.lucene.store.RandomAccessInput; +import org.apache.lucene.util.ArrayUtil; +import org.apache.lucene.util.VectorUtil; + +/** Encodes, ranks, and navigates the centroids used to route vectors into IVF cells. */ +final class Centroids { + private Centroids() {} + + /** + * Keeps centroid representations together so routing can shortlist cheaply and verify accurately. + */ + static final class CentroidCodes { + private static final int TILE = 512; + + final float[][] centroids; + final byte[] coarse; + final int nlist, coarseBytes; + private final int dim, fineStride; + private final FineCodec fine; + private final byte[] fineRecords; + private final ThreadLocal<byte[]> gathered = ThreadLocal.withInitial(() -> new byte[0]); + + CentroidCodes(float[][] centroids, int dim, FineCodec fine) { + this.dim = dim; + this.nlist = centroids.length; + this.centroids = centroids; + this.coarseBytes = Nitrox2.bytesPerVector(dim); + this.coarse = new byte[nlist * coarseBytes]; + this.fine = fine; + this.fineStride = fine == null ? 0 : CodeRecord.length(fine.codeBytes); + this.fineRecords = fine == null ? null : new byte[nlist * fineStride]; + encodeAll(); + } + + void encodeAll() { + for (int c = 0; c < nlist; c++) { + Nitrox2.encode(centroids[c], dim, coarse, c * coarseBytes); + if (fine != null) fine.encode(centroids[c], fineRecords, c * fineStride); + } + } + + void rankCandidates(float[] vector, int[] cands, int count, float[] out) { + if (fine == null) { + for (int i = 0; i < count; i++) out[i] = exactDistance(vector, cands[i]); + return; + } + byte[] flat = gathered.get(); + if (flat.length < count * fineStride) gathered.set(flat = new byte[count * fineStride]); + for (int i = 0; i < count; i++) { + System.arraycopy(fineRecords, cands[i] * fineStride, flat, i * fineStride, fineStride); + } + fine.query(vector, null).score(flat, fineStride, count, out); + for (int i = 0; i < count; i++) out[i] = -out[i]; + } + + static final class Routing { + final int[] cells; + int count, cell2; + float d1, d2; + + Routing(int capacity) { + cells = new int[capacity]; + } + } + + static final class Scratch { + final int[] coarseDist = new int[TILE], verifyCells; + final long[] heap; + final float[] verifyDist; + final byte[] qCode; + + Scratch(int dim, int nlist, int shortlist) { + heap = new long[shortlist]; + verifyCells = new int[shortlist]; + verifyDist = new float[shortlist]; + qCode = new byte[Nitrox2.bytesPerVector(dim)]; + } + } + + void routePacked(float[] vector, int shortlist, int keep, Routing out, Scratch scratch) { + final int want = Math.min(shortlist, nlist); + final long[] heap = scratch.heap; + final int[] cd = scratch.coarseDist; + int n = 0, worst = Integer.MAX_VALUE; + for (int base = 0; base < nlist; base += TILE) { + final int rows = Math.min(TILE, nlist - base); + Kernels.INSTANCE.hamming(scratch.qCode, coarse, base * coarseBytes, rows, cd); + for (int r = 0; r < rows; r++) { + final int dist = cd[r]; + if (n < want) { + heap[n++] = ((long) dist << 32) | (base + r); + if (n == want) { + for (int h = (n >>> 1) - 1; h >= 0; h--) { + CentroidGraph.siftDown(heap, h, n, heap[h], true); + } + worst = (int) (heap[0] >>> 32); + } + } else if (dist < worst) { + CentroidGraph.siftDown(heap, 0, n, ((long) dist << 32) | (base + r), true); + worst = (int) (heap[0] >>> 32); + } + } + } + final int[] cells = scratch.verifyCells; + final float[] dists = scratch.verifyDist; + for (int i = 0; i < n; i++) { + cells[i] = (int) heap[i]; + dists[i] = exactDistance(vector, cells[i]); + } + CentroidGraph.sortByDistance(dists, cells, n); + System.arraycopy(cells, 0, out.cells, 0, out.count = Math.min(keep, n)); + out.cell2 = n > 1 ? cells[1] : -1; + out.d1 = n > 0 ? dists[0] : Float.MAX_VALUE; + out.d2 = n > 1 ? dists[1] : Float.MAX_VALUE; + } + + float exactDistance(float[] vector, int c) { + return -VectorUtil.dotProduct(vector, centroids[c]); + } + + static boolean withinMargin(float d1, float d2, float margin) { + return d2 != Float.MAX_VALUE && d2 - d1 <= (margin - 1f) * Math.abs(d1); + } + } + + /** + * A compact graph over centroid codes that avoids scoring every cell when selecting query probes. + */ + record CentroidGraph( + int nlist, int coarseBytes, int stride, int entry, byte[] nodes, int[][] building) { + static final int M = 16, EF_CONSTRUCTION = 64, EF_MULTIPLIER = 2, MIN_EF = 32; + private static final int ALIGN = 64, ORD_BYTES = 2, LOCK_STRIPES = 512, INSERT_GRAIN = 256; + private static final ThreadLocal<int[]> VISITED = ThreadLocal.withInitial(() -> new int[1]); + + static CentroidGraph build(CentroidCodes codes, int dim) throws IOException { + int nlist = codes.nlist, coarseBytes = codes.coarseBytes; + int stride = (coarseBytes + 2 + M * ORD_BYTES + ALIGN - 1) / ALIGN * ALIGN; + int[][] neighbours = new int[nlist][]; + Arrays.fill(neighbours, new int[0]); + int[] order = new int[nlist]; + for (int i = 0; i < nlist; i++) order[i] = i; + Random random = new Random(0x5DEECE66DL); + for (int i = nlist - 1; i > 0; i--) { + int j = random.nextInt(i + 1), t = order[i]; + order[i] = order[j]; + order[j] = t; + } + int entry = order[0]; + CentroidGraph partial = + new CentroidGraph(nlist, coarseBytes, coarseBytes, entry, codes.coarse, neighbours); + Object[] locks = new Object[LOCK_STRIPES]; + for (int i = 0; i < LOCK_STRIPES; i++) locks[i] = new Object(); + Parallel.RangeTask insertRange = + (from, to) -> { + byte[] code = new byte[coarseBytes]; + int[] visited = new int[nlist]; + for (int idx = from + 1; idx <= to; idx++) { + int node = order[idx]; + Nitrox2.encode(codes.centroids[node], dim, code, 0); + int n = partial.search(code, EF_CONSTRUCTION, visited, null); + int[] kept = neighbours[node] = prune(codes, codes.centroids[node], visited, n); + for (int i = 0; i < kept.length; i++) { + int x = kept[i]; + synchronized (locks[(x * 0x9E3779B9) >>> 1 & (LOCK_STRIPES - 1)]) { + neighbours[x] = link(codes, neighbours[x], x, node, i == 0); + } + } + } + }; + int seed = Math.min(nlist - 1, Math.max(64, M * 4)); + insertRange.run(0, seed); + Parallel.overRange( + nlist - 1 - seed, INSERT_GRAIN, (from, to) -> insertRange.run(seed + from, seed + to)); + connect(neighbours, entry); + byte[] nodes = new byte[nlist * stride]; + for (int c = 0; c < nlist; c++) { + System.arraycopy(codes.coarse, c * coarseBytes, nodes, c * stride, coarseBytes); + int off = c * stride + coarseBytes, deg = Math.min(M, neighbours[c].length); + nodes[off] = (byte) deg; + for (int i = 0; i < deg; i++) { + nodes[off += ORD_BYTES] = (byte) neighbours[c][i]; + nodes[off + 1] = (byte) (neighbours[c][i] >>> 8); + } + } + return new CentroidGraph(nlist, coarseBytes, stride, entry, nodes, null); + } + + private static void connect(int[][] neighbours, int entry) { + int n = neighbours.length; + boolean[] reachable = new boolean[n]; + int[] queue = new int[n]; + for (int guard = 0; guard <= 8; guard++) { + Arrays.fill(reachable, false); + int head = 0, tail = 0; + queue[tail++] = entry; + reachable[entry] = true; + while (head < tail) { + for (int x : neighbours[queue[head++]]) { + if (reachable[x]) continue; + reachable[x] = true; + queue[tail++] = x; + } + } + if (tail == n || guard == 8) return; + for (int c = 0, host = entry; c < n; host = c++) { + if (reachable[c]) continue; + neighbours[c] = appendUnique(neighbours[c], host); + if (neighbours[host].length < M) { + neighbours[host] = appendUnique(neighbours[host], c); + } else { + neighbours[host] = neighbours[host].clone(); + neighbours[host][M - 1] = c; + } + reachable[c] = true; + } + } + } + + private static int[] appendUnique(int[] a, int v) { + for (int x : a) if (x == v) return a; + if (a.length >= M) return a; + int[] out = ArrayUtil.growExact(a, a.length + 1); + out[a.length] = v; + return out; + } + + private static void sortByDistance(float[] dist, int[] ids, int n) { + for (int i = 1; i < n; i++) { + float d = dist[i]; + int c = ids[i], j = i - 1; + for (; j >= 0 && dist[j] > d; j--) { + dist[j + 1] = dist[j]; + ids[j + 1] = ids[j]; + } + dist[j + 1] = d; + ids[j + 1] = c; + } + } + + private static int[] prune(CentroidCodes codes, float[] vec, int[] cand, int n) { + float[] dist = new float[n]; + for (int i = 0; i < n; i++) dist[i] = codes.exactDistance(vec, cand[i]); + sortByDistance(dist, cand, n); + int[] kept = new int[Math.min(M, n)]; + int nKept = 0; + for (int i = 0; i < n && nKept < kept.length; i++) { + boolean diverse = true; + for (int k = 0; k < nKept && diverse; k++) { + diverse = (codes.exactDistance(codes.centroids[cand[i]], kept[k]) < dist[i]) == false; + } + if (diverse) kept[nKept++] = cand[i]; + } + return ArrayUtil.copyOfSubArray(kept, 0, nKept); + } + + private static int[] link(CentroidCodes codes, int[] cur, int x, int node, boolean mustLink) { + if (cur.length < M) return appendUnique(cur, node); + for (int y : cur) if (y == node) return cur; + float[] xVec = codes.centroids[x]; + int worst = -1; + float worstD = Float.NEGATIVE_INFINITY; + for (int i = 0; i < cur.length; i++) { + float d = codes.exactDistance(xVec, cur[i]); + if (d > worstD) { + worstD = d; + worst = i; + } + } + if ((mustLink || codes.exactDistance(xVec, node) < worstD) == false) return cur; + int[] out = cur.clone(); + out[worst] = node; + return out; + } + + int search(byte[] qCode, int ef, int[] out, int[] outDist) { + int[] visited = VISITED.get(); + if (visited.length <= nlist) VISITED.set(visited = new int[nlist + 1]); + int gen = ++visited[0], nOut = 0, frontierN = 0, bestN = 0; + int cap = Math.min(Math.max(ef, 1), nlist); + long[] frontier = new long[Math.min(nlist, Math.max(64, cap * 4))], best = new long[cap]; + int[] fanOffsets = new int[M], fanDist = new int[M + 1]; + visited[entry + 1] = gen; + fanOffsets[0] = entry * stride; + for (int fan = 1; ; ) { + int firstOut = nOut; + for (int i = 0; i < fan && nOut < out.length; i++) out[nOut++] = fanOffsets[i] / stride; + if (fan > 0) Kernels.INSTANCE.hammingAt(qCode, nodes, fanOffsets, fan, fanDist); + if (outDist != null) System.arraycopy(fanDist, 0, outDist, firstOut, nOut - firstOut); + for (int i = 0; i < fan; i++) { + if (bestN == cap && fanDist[i] >= (int) (best[0] >>> 32)) continue; + long e = ((long) fanDist[i] << 32) | (fanOffsets[i] / stride); + if (bestN < cap) siftUp(best, bestN++, e, true); + else siftDown(best, 0, bestN, e, true); + if (frontierN < frontier.length) siftUp(frontier, frontierN++, e, false); + } + if (frontierN == 0) return nOut; + long top = frontier[0]; + siftDown(frontier, 0, --frontierN, frontier[frontierN], false); + if (bestN == cap && (int) (top >>> 32) > (int) (best[0] >>> 32)) return nOut; + int node = (int) top; + int[] adj = building == null ? null : building[node]; + int degOff = node * stride + coarseBytes; + int deg = + adj != null ? adj.length : (nodes[degOff] & 0xFF) | (nodes[degOff + 1] & 0xFF) << 8; + fan = 0; + for (int i = 0, off = degOff + 2; i < deg; i++, off += ORD_BYTES) { Review Comment: I wonder if this loop would be more readable if you had separate logic for the case when building != null and when building == null... ########## lucene/sandbox/src/java/org/apache/lucene/sandbox/codecs/segmentivf/Centroids.java: ########## @@ -0,0 +1,375 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.lucene.sandbox.codecs.segmentivf; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Random; +import org.apache.lucene.sandbox.codecs.segmentivf.Clustering.Parallel; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.CodeRecord; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.FineCodec; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.Nitrox2; +import org.apache.lucene.store.ByteArrayDataInput; +import org.apache.lucene.store.IndexOutput; +import org.apache.lucene.store.RandomAccessInput; +import org.apache.lucene.util.ArrayUtil; +import org.apache.lucene.util.VectorUtil; + +/** Encodes, ranks, and navigates the centroids used to route vectors into IVF cells. */ +final class Centroids { + private Centroids() {} + + /** + * Keeps centroid representations together so routing can shortlist cheaply and verify accurately. + */ + static final class CentroidCodes { + private static final int TILE = 512; + + final float[][] centroids; + final byte[] coarse; + final int nlist, coarseBytes; + private final int dim, fineStride; + private final FineCodec fine; + private final byte[] fineRecords; + private final ThreadLocal<byte[]> gathered = ThreadLocal.withInitial(() -> new byte[0]); + + CentroidCodes(float[][] centroids, int dim, FineCodec fine) { + this.dim = dim; + this.nlist = centroids.length; + this.centroids = centroids; + this.coarseBytes = Nitrox2.bytesPerVector(dim); + this.coarse = new byte[nlist * coarseBytes]; + this.fine = fine; + this.fineStride = fine == null ? 0 : CodeRecord.length(fine.codeBytes); + this.fineRecords = fine == null ? null : new byte[nlist * fineStride]; + encodeAll(); + } + + void encodeAll() { + for (int c = 0; c < nlist; c++) { + Nitrox2.encode(centroids[c], dim, coarse, c * coarseBytes); + if (fine != null) fine.encode(centroids[c], fineRecords, c * fineStride); + } + } + + void rankCandidates(float[] vector, int[] cands, int count, float[] out) { + if (fine == null) { + for (int i = 0; i < count; i++) out[i] = exactDistance(vector, cands[i]); + return; + } + byte[] flat = gathered.get(); + if (flat.length < count * fineStride) gathered.set(flat = new byte[count * fineStride]); + for (int i = 0; i < count; i++) { + System.arraycopy(fineRecords, cands[i] * fineStride, flat, i * fineStride, fineStride); + } + fine.query(vector, null).score(flat, fineStride, count, out); + for (int i = 0; i < count; i++) out[i] = -out[i]; + } + + static final class Routing { + final int[] cells; + int count, cell2; + float d1, d2; + + Routing(int capacity) { + cells = new int[capacity]; + } + } + + static final class Scratch { + final int[] coarseDist = new int[TILE], verifyCells; + final long[] heap; + final float[] verifyDist; + final byte[] qCode; + + Scratch(int dim, int nlist, int shortlist) { + heap = new long[shortlist]; + verifyCells = new int[shortlist]; + verifyDist = new float[shortlist]; + qCode = new byte[Nitrox2.bytesPerVector(dim)]; + } + } + + void routePacked(float[] vector, int shortlist, int keep, Routing out, Scratch scratch) { + final int want = Math.min(shortlist, nlist); + final long[] heap = scratch.heap; + final int[] cd = scratch.coarseDist; + int n = 0, worst = Integer.MAX_VALUE; + for (int base = 0; base < nlist; base += TILE) { + final int rows = Math.min(TILE, nlist - base); + Kernels.INSTANCE.hamming(scratch.qCode, coarse, base * coarseBytes, rows, cd); + for (int r = 0; r < rows; r++) { + final int dist = cd[r]; + if (n < want) { + heap[n++] = ((long) dist << 32) | (base + r); + if (n == want) { + for (int h = (n >>> 1) - 1; h >= 0; h--) { + CentroidGraph.siftDown(heap, h, n, heap[h], true); + } + worst = (int) (heap[0] >>> 32); + } + } else if (dist < worst) { + CentroidGraph.siftDown(heap, 0, n, ((long) dist << 32) | (base + r), true); + worst = (int) (heap[0] >>> 32); + } + } + } + final int[] cells = scratch.verifyCells; + final float[] dists = scratch.verifyDist; + for (int i = 0; i < n; i++) { + cells[i] = (int) heap[i]; + dists[i] = exactDistance(vector, cells[i]); + } + CentroidGraph.sortByDistance(dists, cells, n); + System.arraycopy(cells, 0, out.cells, 0, out.count = Math.min(keep, n)); + out.cell2 = n > 1 ? cells[1] : -1; + out.d1 = n > 0 ? dists[0] : Float.MAX_VALUE; + out.d2 = n > 1 ? dists[1] : Float.MAX_VALUE; + } + + float exactDistance(float[] vector, int c) { + return -VectorUtil.dotProduct(vector, centroids[c]); + } + + static boolean withinMargin(float d1, float d2, float margin) { + return d2 != Float.MAX_VALUE && d2 - d1 <= (margin - 1f) * Math.abs(d1); + } + } + + /** + * A compact graph over centroid codes that avoids scoring every cell when selecting query probes. + */ + record CentroidGraph( + int nlist, int coarseBytes, int stride, int entry, byte[] nodes, int[][] building) { + static final int M = 16, EF_CONSTRUCTION = 64, EF_MULTIPLIER = 2, MIN_EF = 32; + private static final int ALIGN = 64, ORD_BYTES = 2, LOCK_STRIPES = 512, INSERT_GRAIN = 256; + private static final ThreadLocal<int[]> VISITED = ThreadLocal.withInitial(() -> new int[1]); + + static CentroidGraph build(CentroidCodes codes, int dim) throws IOException { + int nlist = codes.nlist, coarseBytes = codes.coarseBytes; + int stride = (coarseBytes + 2 + M * ORD_BYTES + ALIGN - 1) / ALIGN * ALIGN; + int[][] neighbours = new int[nlist][]; + Arrays.fill(neighbours, new int[0]); + int[] order = new int[nlist]; + for (int i = 0; i < nlist; i++) order[i] = i; + Random random = new Random(0x5DEECE66DL); + for (int i = nlist - 1; i > 0; i--) { + int j = random.nextInt(i + 1), t = order[i]; + order[i] = order[j]; + order[j] = t; + } + int entry = order[0]; + CentroidGraph partial = + new CentroidGraph(nlist, coarseBytes, coarseBytes, entry, codes.coarse, neighbours); + Object[] locks = new Object[LOCK_STRIPES]; + for (int i = 0; i < LOCK_STRIPES; i++) locks[i] = new Object(); + Parallel.RangeTask insertRange = + (from, to) -> { + byte[] code = new byte[coarseBytes]; + int[] visited = new int[nlist]; + for (int idx = from + 1; idx <= to; idx++) { + int node = order[idx]; + Nitrox2.encode(codes.centroids[node], dim, code, 0); + int n = partial.search(code, EF_CONSTRUCTION, visited, null); + int[] kept = neighbours[node] = prune(codes, codes.centroids[node], visited, n); + for (int i = 0; i < kept.length; i++) { + int x = kept[i]; + synchronized (locks[(x * 0x9E3779B9) >>> 1 & (LOCK_STRIPES - 1)]) { + neighbours[x] = link(codes, neighbours[x], x, node, i == 0); + } + } + } + }; + int seed = Math.min(nlist - 1, Math.max(64, M * 4)); + insertRange.run(0, seed); + Parallel.overRange( + nlist - 1 - seed, INSERT_GRAIN, (from, to) -> insertRange.run(seed + from, seed + to)); + connect(neighbours, entry); + byte[] nodes = new byte[nlist * stride]; + for (int c = 0; c < nlist; c++) { + System.arraycopy(codes.coarse, c * coarseBytes, nodes, c * stride, coarseBytes); + int off = c * stride + coarseBytes, deg = Math.min(M, neighbours[c].length); + nodes[off] = (byte) deg; + for (int i = 0; i < deg; i++) { + nodes[off += ORD_BYTES] = (byte) neighbours[c][i]; + nodes[off + 1] = (byte) (neighbours[c][i] >>> 8); + } + } + return new CentroidGraph(nlist, coarseBytes, stride, entry, nodes, null); + } + + private static void connect(int[][] neighbours, int entry) { + int n = neighbours.length; + boolean[] reachable = new boolean[n]; + int[] queue = new int[n]; + for (int guard = 0; guard <= 8; guard++) { + Arrays.fill(reachable, false); + int head = 0, tail = 0; + queue[tail++] = entry; + reachable[entry] = true; + while (head < tail) { + for (int x : neighbours[queue[head++]]) { + if (reachable[x]) continue; + reachable[x] = true; + queue[tail++] = x; + } + } + if (tail == n || guard == 8) return; + for (int c = 0, host = entry; c < n; host = c++) { + if (reachable[c]) continue; + neighbours[c] = appendUnique(neighbours[c], host); + if (neighbours[host].length < M) { + neighbours[host] = appendUnique(neighbours[host], c); + } else { + neighbours[host] = neighbours[host].clone(); + neighbours[host][M - 1] = c; + } + reachable[c] = true; + } + } + } + + private static int[] appendUnique(int[] a, int v) { + for (int x : a) if (x == v) return a; + if (a.length >= M) return a; + int[] out = ArrayUtil.growExact(a, a.length + 1); + out[a.length] = v; + return out; + } + + private static void sortByDistance(float[] dist, int[] ids, int n) { + for (int i = 1; i < n; i++) { + float d = dist[i]; + int c = ids[i], j = i - 1; + for (; j >= 0 && dist[j] > d; j--) { + dist[j + 1] = dist[j]; + ids[j + 1] = ids[j]; + } + dist[j + 1] = d; + ids[j + 1] = c; + } + } + + private static int[] prune(CentroidCodes codes, float[] vec, int[] cand, int n) { + float[] dist = new float[n]; + for (int i = 0; i < n; i++) dist[i] = codes.exactDistance(vec, cand[i]); + sortByDistance(dist, cand, n); + int[] kept = new int[Math.min(M, n)]; + int nKept = 0; + for (int i = 0; i < n && nKept < kept.length; i++) { + boolean diverse = true; + for (int k = 0; k < nKept && diverse; k++) { + diverse = (codes.exactDistance(codes.centroids[cand[i]], kept[k]) < dist[i]) == false; + } + if (diverse) kept[nKept++] = cand[i]; + } + return ArrayUtil.copyOfSubArray(kept, 0, nKept); + } + + private static int[] link(CentroidCodes codes, int[] cur, int x, int node, boolean mustLink) { + if (cur.length < M) return appendUnique(cur, node); + for (int y : cur) if (y == node) return cur; + float[] xVec = codes.centroids[x]; + int worst = -1; + float worstD = Float.NEGATIVE_INFINITY; + for (int i = 0; i < cur.length; i++) { + float d = codes.exactDistance(xVec, cur[i]); + if (d > worstD) { + worstD = d; + worst = i; + } + } + if ((mustLink || codes.exactDistance(xVec, node) < worstD) == false) return cur; + int[] out = cur.clone(); + out[worst] = node; + return out; + } + + int search(byte[] qCode, int ef, int[] out, int[] outDist) { Review Comment: I found this method quite difficult to read. Not sure if more descriptive variable names would help. Or a descriptive comment of the algorithm. Or both? I don't know if it's just me or if others will also snuggle through this, but you might consider ways to make these core algorithms a little more readable by humans. ########## lucene/sandbox/src/java/org/apache/lucene/sandbox/codecs/segmentivf/SegmentIVFVectorsReader.java: ########## @@ -0,0 +1,1241 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.lucene.sandbox.codecs.segmentivf; + +import static org.apache.lucene.codecs.CodecUtil.checkIndexHeader; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DATA_CODEC_NAME; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DATA_EXTENSION; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DEFAULT_PROBE_MARGIN; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DIRECT_MONOTONIC_BLOCK_SHIFT; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.META_CODEC_NAME; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.META_EXTENSION; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.VERSION_CURRENT; +import static org.apache.lucene.search.DocIdSetIterator.NO_MORE_DOCS; +import static org.apache.lucene.util.packed.DirectMonotonicReader.loadMeta; + +import java.io.Closeable; +import java.io.IOException; +import java.io.UncheckedIOException; +import java.lang.foreign.AddressLayout; +import java.lang.foreign.Arena; +import java.lang.foreign.FunctionDescriptor; +import java.lang.foreign.Linker; +import java.lang.foreign.MemorySegment; +import java.lang.foreign.ValueLayout; +import java.lang.invoke.MethodHandle; +import java.lang.invoke.VarHandle; +import java.nio.ByteOrder; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Arrays; +import java.util.HashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicLong; +import org.apache.lucene.codecs.CodecUtil; +import org.apache.lucene.codecs.KnnVectorsReader; +import org.apache.lucene.index.ByteVectorValues; +import org.apache.lucene.index.CorruptIndexException; +import org.apache.lucene.index.FieldInfo; +import org.apache.lucene.index.Float16VectorValues; +import org.apache.lucene.index.FloatVectorValues; +import org.apache.lucene.index.IndexFileNames; +import org.apache.lucene.index.MergePolicy; +import org.apache.lucene.index.SegmentReadState; +import org.apache.lucene.index.VectorSimilarityFunction; +import org.apache.lucene.sandbox.codecs.segmentivf.Centroids.CentroidCodes; +import org.apache.lucene.sandbox.codecs.segmentivf.Centroids.CentroidGraph; +import org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.FineTier; +import org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.SearchStrategy; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.CodeRecord; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.FineCodec; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.HadamardRotation; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.Nitrox2; +import org.apache.lucene.search.AcceptDocs; +import org.apache.lucene.search.DocIdSetIterator; +import org.apache.lucene.search.KnnCollector; +import org.apache.lucene.search.VectorScorer; +import org.apache.lucene.store.ChecksumIndexInput; +import org.apache.lucene.store.FSDirectory; +import org.apache.lucene.store.FilterDirectory; +import org.apache.lucene.store.IOContext; +import org.apache.lucene.store.IndexInput; +import org.apache.lucene.store.MemorySegmentAccessInput; +import org.apache.lucene.store.RandomAccessInput; +import org.apache.lucene.util.ArrayUtil; +import org.apache.lucene.util.BitSet; +import org.apache.lucene.util.BitUtil; +import org.apache.lucene.util.Bits; +import org.apache.lucene.util.IOUtils; +import org.apache.lucene.util.NumericUtils; +import org.apache.lucene.util.VectorUtil; +import org.apache.lucene.util.packed.DirectMonotonicReader; + +/** + * Searches SegmentIVF fields by probing cells, scanning coarse codes, and reranking a shortlist. + * + * <p>The coarse scan is intentionally bandwidth-oriented: compact Nitrox2 rows are processed with + * XOR and popcount, then only a bounded set of survivors reaches the more expensive fine scorer. + * Filtered search chooses between scanning selected cells and visiting accepted documents directly. + */ +final class SegmentIVFVectorsReader extends KnnVectorsReader { + /** Coarse candidates fine-reranked per requested neighbor, and never fewer than MIN_RERANK. */ + static final int RERANK_PER_K = 7, MIN_RERANK = 100; + + /** Returns how many coarse candidates to fine-rerank for a top-{@code k} search. */ + static long rerankCount(int k) { + return Math.max(MIN_RERANK, (long) RERANK_PER_K * k); + } + + private static final int ADMIT_BLOCK = 256, FILTERED_PROBE_MULTIPLIER = 8; + private static final int VERIFY_MIN = 64, VERIFY_MULTIPLIER = 2; + private static final Kernels K = Kernels.INSTANCE; + private static final ValueLayout.OfInt INT_LE = + ValueLayout.JAVA_INT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN); + + /** + * A segment's deduplicated coarse shortlist: record slots and their coarse distances, plus the + * search that gathered them, whose encoded query the rerank reuses. + */ + record Candidates(int[] slots, int[] distances, Field.Search search) { + static final Candidates EMPTY = new Candidates(new int[0], new int[0], null); + } + + private final Map<String, Field> fields = new HashMap<>(); + private final IndexInput data; + private final Uring uring; + private final Arena pinned = Arena.ofShared(); + private final AtomicLong pinnedBytes = new AtomicLong(); + private long fineBytes; + private boolean closed; + private final String segment; + private final int segmentMaxDoc; + + SegmentIVFVectorsReader(SegmentReadState state) throws IOException { + String name = segment = state.segmentInfo.name, sfx = state.segmentSuffix; + segmentMaxDoc = state.segmentInfo.maxDoc(); + byte[] id = state.segmentInfo.getId(); + String metaName = IndexFileNames.segmentFileName(name, sfx, META_EXTENSION); + try (ChecksumIndexInput meta = state.directory.openChecksumInput(metaName)) { + Throwable prior = null; + try { + checkIndexHeader(meta, META_CODEC_NAME, VERSION_CURRENT, VERSION_CURRENT, id, sfx); + for (int number = meta.readInt(); number != -1; number = meta.readInt()) { + FieldInfo info = state.fieldInfos.fieldInfo(number); + if (info == null) throw new CorruptIndexException("invalid field number " + number, meta); + fields.put(info.name, new Field(meta, info)); + } + } catch (Throwable t) { + prior = t; + } finally { + CodecUtil.checkFooter(meta, prior); + } + } + String dataName = IndexFileNames.segmentFileName(name, sfx, DATA_EXTENSION); + data = state.directory.openInput(dataName, state.context); + try { + checkIndexHeader(data, DATA_CODEC_NAME, VERSION_CURRENT, VERSION_CURRENT, id, sfx); + CodecUtil.retrieveChecksum(data); + } catch (Throwable t) { + IOUtils.closeWhileSuppressingExceptions(t, data); + throw t; + } + uring = openUring(state, dataName); + for (Field field : fields.values()) fineBytes += field.sections[2] - field.sections[1]; + Uring.addFine(fineBytes); + if (state.context.context() != IOContext.Context.MERGE) { + for (Field field : fields.values()) field.maybePin(); + } + } + + /** Opens batched fine-record reads for a data file of a file-system directory, if possible. */ + private static Uring openUring(SegmentReadState state, String name) { + if (FilterDirectory.unwrap(state.directory) instanceof FSDirectory fs) { + try { + return Uring.open(fs.getDirectory().resolve(name)); + } catch (IOException _) { + // fine records are read through the mapped input instead + } + } + return null; + } + + final class Field { + final VectorSimilarityFunction similarity; + final FineTier fineTier; + final int dim, nlist, count, nprobe, spillBits, recordLen, coarseBytes, docIdOffset; + final long rotationSeed; + final long[] sections; + final DirectMonotonicReader.Meta postingOffsets; + final HadamardRotation rotation; + final FineCodec fine; + + RandomAccessInput records, slotDocs; + MemorySegment recordsSeg; // the whole mapped records section, or null + MemorySegmentAccessInput coarseAccess; + MemorySegment coarseSeg; + // Off-heap copies published by the pinning thread; searches use the mapped sections until then. + volatile MemorySegment pinnedCoarse, pinnedSlotDocs; + private final AtomicBoolean pinning = new AtomicBoolean(); + private volatile long lastPinAttempt; + int[] cellStart, ordToSlot, ordToDoc, allCells, primaryCells; + float[][] centroids; + private volatile CentroidCodes codes; + CentroidGraph graph; + + Field(ChecksumIndexInput meta, FieldInfo info) throws IOException { + similarity = info.getVectorSimilarityFunction(); + fineTier = FineTier.values()[meta.readByte()]; + dim = meta.readVInt(); + nlist = meta.readVInt(); + count = meta.readVInt(); + rotationSeed = meta.readLong(); + nprobe = meta.readVInt(); + spillBits = meta.readVInt(); + // Sections: centroids, records, coarse, graph, ordToSlot, slotDoc, posting offsets. + sections = new long[7]; + for (int s = 0; s < sections.length; s++) sections[s] = meta.readVLong(); + postingOffsets = nlist == 0 ? null : loadMeta(meta, nlist + 1, DIRECT_MONOTONIC_BLOCK_SHIFT); + if (dim != info.getVectorDimension()) throw new CorruptIndexException("dimension", meta); + rotation = HadamardRotation.create(dim, rotationSeed); + fine = new FineCodec(fineTier, dim); + recordLen = CodeRecord.length(fine.codeBytes); + coarseBytes = Nitrox2.bytesPerVector(dim); + docIdOffset = fine.codeBytes; + } + + private RandomAccessInput section(int s) throws IOException { + return data.randomAccessSlice(sections[s], sections[s + 1] - sections[s]); + } + + synchronized Field open() throws IOException { + if (records != null) return this; + if (section(2) instanceof MemorySegmentAccessInput in) { + coarseAccess = in; + coarseSeg = segmentOrNull(in, 0, in.length()); + } + cellStart = new int[nlist + 1]; + if (nlist > 0) { + long postings = sections[sections.length - 1]; + RandomAccessInput tail = data.randomAccessSlice(postings, data.length() - postings); + var offsets = DirectMonotonicReader.getInstance(postingOffsets, tail); + for (int c = 0; c <= nlist; c++) cellStart[c] = (int) (offsets.get(c) / Integer.BYTES); + } + slotDocs = section(5); + IndexInput all = data.clone(); + centroids = new float[nlist][dim]; + all.seek(sections[0]); + for (float[] centroid : centroids) all.readFloats(centroid, 0, dim); + records = section(1); + if (records instanceof MemorySegmentAccessInput in) { + recordsSeg = segmentOrNull(in, 0, in.length()); + } + return this; + } + + /** + * Queues the off-heap copy of the coarse and slot-to-document sections when the pin budget has + * room, retrying at most once per second: a merged segment is pinned once its sources close. + * Searches never wait for the copy. + */ + void maybePin() { + if (count == 0 || pinnedCoarse != null || pinning.get()) return; + long now = System.nanoTime(); + if (lastPinAttempt != 0 && now - lastPinAttempt < 1_000_000_000L) return; + if (pinning.getAndSet(true)) return; + lastPinAttempt = now == 0 ? 1 : now; + long coarseLength = sections[3] - sections[2]; + long docsLength = sections[6] - sections[5]; + if (Uring.reservePinned(coarseLength + docsLength) == false) { + pinning.set(false); + return; + } + pinnedBytes.addAndGet(coarseLength + docsLength); + Uring.PINNER.execute( + () -> { + try { + if (docsLength > 0) pinnedSlotDocs = pin(sections[5], docsLength); + pinnedCoarse = pin(sections[2], coarseLength); + } catch (IOException | RuntimeException _) { + // the reader closed mid-copy; its reservation is released on close + } + }); + } + + /** Copies a section off-heap, where the page cache cannot evict it. */ + private MemorySegment pin(long offset, long length) throws IOException { + MemorySegment copy = pinned.allocate(length, 64); // cache-line aligned for the SIMD scans + IndexInput in = data.clone(); + in.seek(offset); + byte[] chunk = new byte[1 << 20]; + for (long at = 0; at < length; at += chunk.length) { + int n = (int) Math.min(chunk.length, length - at); + in.readBytes(chunk, 0, n); + MemorySegment.copy(chunk, 0, copy, ValueLayout.JAVA_BYTE, at, n); + } + return copy; + } + + private int docAt(RandomAccessInput docs, int slot) throws IOException { + MemorySegment pinned = pinnedSlotDocs; + return pinned != null ? pinned.getAtIndex(INT_LE, slot) : docs.readInt((long) slot * 4); + } + + private synchronized void loadOrdToSlot() throws IOException { + if (ordToSlot != null) return; + int[] slots = new int[count]; + IndexInput all = data.clone(); + all.seek(sections[4]); + all.readInts(slots, 0, count); + ordToSlot = slots; + } + + private synchronized void loadOrdinalMappings() throws IOException { + if (ordToDoc != null) return; + loadOrdToSlot(); + int[] docs = new int[count]; + for (int ord = 0; ord < count; ord++) docs[ord] = docAt(slotDocs, ordToSlot[ord]); + ordToDoc = docs; + } + + /** + * Returns a mapped slice, or null when no single mapping covers it. The slice is rebased as a + * plain native segment scoped to this reader, so the SIMD kernels see one segment type whether + * or not a section is pinned: once they have seen both, every coarse scan runs ~16% slower. + */ + @SuppressWarnings("restricted") + private MemorySegment segmentOrNull(MemorySegmentAccessInput in, long offset, long length) { + if (length == 0) return null; + try { + MemorySegment mapped = in.segmentSliceOrNull(offset, length); + if (mapped == null) return null; + return MemorySegment.ofAddress(mapped.address()).reinterpret(length, pinned, null); + } catch (IOException _) { + return null; + } + } + + /** + * Returns the mapping holding one cell run: the whole coarse section when mappable, otherwise a + * slice rebased to the run, or null when neither is mapped. Offsets start at {@code runBase}. + */ + private MemorySegment coarseRun(int slotBase, int rows) { + if (coarseSeg != null || coarseAccess == null) return coarseSeg; + return segmentOrNull(coarseAccess, (long) slotBase * coarseBytes, (long) rows * coarseBytes); + } + + private long runBase(int slotBase) { + return coarseSeg == null ? 0 : (long) slotBase * coarseBytes; + } + + private synchronized void loadCodes() throws IOException { + if (codes != null) return; + allCells = new int[nlist]; + for (int c = 0; c < nlist; c++) allCells[c] = c; + long graphLength = sections[4] - sections[3]; + graph = graphLength == 0 ? null : CentroidGraph.read(section(3), dim, graphLength); + codes = new CentroidCodes(centroids, dim, fine); + } + + synchronized int cellOf(int ord) throws IOException { + loadOrdToSlot(); + if (primaryCells == null) { + primaryCells = new int[count]; + for (int o = 0; o < count; o++) { + primaryCells[o] = records.readInt((long) ordToSlot[o] * recordLen + docIdOffset + 4); + } + } + return primaryCells[ord]; + } + + /** + * Per-query state for cell selection, coarse admission, deduplication, and fine reranking. Each + * search opens its own slices, which shadow the field's: positional reads through a shared + * slice are not thread-safe on every directory implementation. + */ + final class Search { + final RandomAccessInput records = section(1), coarse = section(2); + final RandomAccessInput slotDocs = section(5); + // One snapshot per query, so a copy published mid-query cannot mix offsets. + final MemorySegment pinnedCoarse = Field.this.pinnedCoarse; + final float[] rotated; + final byte[] qCode; + final FineCodec.Query fine; + final KnnCollector collector; + final Scratch scratch = Scratch.LOCAL.get(); + final int bins = coarseBytes * 8 + 2, shortlist, pool; + final int[] histogram = scratch.histogram = ArrayUtil.growNoCopy(scratch.histogram, bins); + int size, admitted, threshold = bins - 1; + Candidates gathered; + + Search(float[] target, KnnCollector collector, int k) throws IOException { + maybePin(); + this.collector = collector; + shortlist = (int) Math.min(count, rerankCount(k)); + pool = Math.multiplyExact(shortlist, 1 + spillBits); + scratch.reserve(shortlist); + Arrays.fill(histogram, 0, bins, 0); + if (codes == null) loadCodes(); + rotated = new float[dim]; + qCode = new byte[coarseBytes]; + rotation.rotate(VectorUtil.l2normalize(ArrayUtil.copyOfSubArray(target, 0, dim)), rotated); + Nitrox2.encode(rotated, dim, qCode, 0); + fine = Field.this.fine.query(rotated, similarity); + } + + /** A rerank of {@code from}'s shortlist on this thread, reusing its encoded query. */ + private Search(Search from, KnnCollector collector) throws IOException { + this.collector = collector; + shortlist = from.shortlist; + pool = from.pool; + rotated = from.rotated; + qCode = from.qCode; + fine = from.fine; + } + + /** Fine-reranks {@code slots}, taken from this search's shortlist, into {@code collector}. */ + void rerankInto(int[] slots, KnnCollector collector) throws IOException { + new Search(this, collector).rerank(slots, slots.length); + } + + void run(AcceptDocs acceptDocs) throws IOException { + var strategy = collector.getSearchStrategy(); + if (strategy instanceof SearchStrategy s) run(s.numProbes, s.probeMargin, acceptDocs); + else run(nprobe, DEFAULT_PROBE_MARGIN, acceptDocs); + } + + private void run(int probe, float margin, AcceptDocs acceptDocs) throws IOException { + probe = Math.min(probe, nlist); + Bits accept = acceptDocs == null ? null : acceptDocs.bits(); + if (accept instanceof BitSet filter) { + int cost = acceptDocs.cost(); + double parity = + Math.sqrt((double) shortlist * cellStart[nlist] * coarseBytes / recordLen); + if (cost > (int) Math.max(shortlist, Math.min(Integer.MAX_VALUE, parity))) { + filteredScan(selectCells(probe, 1f), filter, cost, probe); + return; + } + boolean dense = count == segmentMaxDoc; + if (dense) loadOrdToSlot(); + else loadOrdinalMappings(); + int[] slots = new int[64]; + int n = 0; + DocIdSetIterator accepted = acceptDocs.iterator(); + for (int doc = accepted.nextDoc(); doc != NO_MORE_DOCS; doc = accepted.nextDoc()) { + int ord = dense ? doc : Arrays.binarySearch(ordToDoc, doc); + if (ord < 0) continue; + slots = ArrayUtil.grow(slots, n + 1); + slots[n++] = ordToSlot[ord]; + } + if (collector == null) admitAll(slots, n); + else rerank(slots, n); + } else { + // A segment no larger than the rerank pool scans every slot: its mostly empty cells would + // spend probes on cells holding nothing. + scan(cellStart[nlist] <= pool ? allCells : selectCells(probe, margin), accept); + } + } + + private int[] selectCells(int probe, float margin) { + int[] candidates = allCells; + int got = nlist; + if (graph != null) { + candidates = scratch.candidates = ArrayUtil.growNoCopy(scratch.candidates, nlist); + int[] coarse = scratch.coarse = ArrayUtil.growNoCopy(scratch.coarse, nlist); + int ef = Math.max(CentroidGraph.MIN_EF, probe * CentroidGraph.EF_MULTIPLIER); + got = graph.search(qCode, ef, candidates, coarse); + int cap = Math.max(VERIFY_MIN, probe * VERIFY_MULTIPLIER); + if (cap < got) { + int[] counts = new int[bins]; Review Comment: As a general bit of feedback, I think some comments would go a long way, especially in these more core algorithm bits of your work. For example, it took me a minute to realize this is a counting sort. A comment describing the fact that you use a counting sort when you get back more scored centroids than the cap would go a long way for readers. ########## lucene/sandbox/src/java/org/apache/lucene/sandbox/codecs/segmentivf/SegmentIVFVectorsReader.java: ########## @@ -0,0 +1,1241 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.lucene.sandbox.codecs.segmentivf; + +import static org.apache.lucene.codecs.CodecUtil.checkIndexHeader; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DATA_CODEC_NAME; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DATA_EXTENSION; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DEFAULT_PROBE_MARGIN; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.DIRECT_MONOTONIC_BLOCK_SHIFT; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.META_CODEC_NAME; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.META_EXTENSION; +import static org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.VERSION_CURRENT; +import static org.apache.lucene.search.DocIdSetIterator.NO_MORE_DOCS; +import static org.apache.lucene.util.packed.DirectMonotonicReader.loadMeta; + +import java.io.Closeable; +import java.io.IOException; +import java.io.UncheckedIOException; +import java.lang.foreign.AddressLayout; +import java.lang.foreign.Arena; +import java.lang.foreign.FunctionDescriptor; +import java.lang.foreign.Linker; +import java.lang.foreign.MemorySegment; +import java.lang.foreign.ValueLayout; +import java.lang.invoke.MethodHandle; +import java.lang.invoke.VarHandle; +import java.nio.ByteOrder; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Arrays; +import java.util.HashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicLong; +import org.apache.lucene.codecs.CodecUtil; +import org.apache.lucene.codecs.KnnVectorsReader; +import org.apache.lucene.index.ByteVectorValues; +import org.apache.lucene.index.CorruptIndexException; +import org.apache.lucene.index.FieldInfo; +import org.apache.lucene.index.Float16VectorValues; +import org.apache.lucene.index.FloatVectorValues; +import org.apache.lucene.index.IndexFileNames; +import org.apache.lucene.index.MergePolicy; +import org.apache.lucene.index.SegmentReadState; +import org.apache.lucene.index.VectorSimilarityFunction; +import org.apache.lucene.sandbox.codecs.segmentivf.Centroids.CentroidCodes; +import org.apache.lucene.sandbox.codecs.segmentivf.Centroids.CentroidGraph; +import org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.FineTier; +import org.apache.lucene.sandbox.codecs.segmentivf.SegmentIVFVectorsFormat.SearchStrategy; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.CodeRecord; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.FineCodec; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.HadamardRotation; +import org.apache.lucene.sandbox.codecs.segmentivf.Tiers.Nitrox2; +import org.apache.lucene.search.AcceptDocs; +import org.apache.lucene.search.DocIdSetIterator; +import org.apache.lucene.search.KnnCollector; +import org.apache.lucene.search.VectorScorer; +import org.apache.lucene.store.ChecksumIndexInput; +import org.apache.lucene.store.FSDirectory; +import org.apache.lucene.store.FilterDirectory; +import org.apache.lucene.store.IOContext; +import org.apache.lucene.store.IndexInput; +import org.apache.lucene.store.MemorySegmentAccessInput; +import org.apache.lucene.store.RandomAccessInput; +import org.apache.lucene.util.ArrayUtil; +import org.apache.lucene.util.BitSet; +import org.apache.lucene.util.BitUtil; +import org.apache.lucene.util.Bits; +import org.apache.lucene.util.IOUtils; +import org.apache.lucene.util.NumericUtils; +import org.apache.lucene.util.VectorUtil; +import org.apache.lucene.util.packed.DirectMonotonicReader; + +/** + * Searches SegmentIVF fields by probing cells, scanning coarse codes, and reranking a shortlist. + * + * <p>The coarse scan is intentionally bandwidth-oriented: compact Nitrox2 rows are processed with + * XOR and popcount, then only a bounded set of survivors reaches the more expensive fine scorer. + * Filtered search chooses between scanning selected cells and visiting accepted documents directly. + */ +final class SegmentIVFVectorsReader extends KnnVectorsReader { + /** Coarse candidates fine-reranked per requested neighbor, and never fewer than MIN_RERANK. */ + static final int RERANK_PER_K = 7, MIN_RERANK = 100; + + /** Returns how many coarse candidates to fine-rerank for a top-{@code k} search. */ + static long rerankCount(int k) { + return Math.max(MIN_RERANK, (long) RERANK_PER_K * k); + } + + private static final int ADMIT_BLOCK = 256, FILTERED_PROBE_MULTIPLIER = 8; + private static final int VERIFY_MIN = 64, VERIFY_MULTIPLIER = 2; + private static final Kernels K = Kernels.INSTANCE; + private static final ValueLayout.OfInt INT_LE = + ValueLayout.JAVA_INT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN); + + /** + * A segment's deduplicated coarse shortlist: record slots and their coarse distances, plus the + * search that gathered them, whose encoded query the rerank reuses. + */ + record Candidates(int[] slots, int[] distances, Field.Search search) { + static final Candidates EMPTY = new Candidates(new int[0], new int[0], null); + } + + private final Map<String, Field> fields = new HashMap<>(); + private final IndexInput data; + private final Uring uring; + private final Arena pinned = Arena.ofShared(); + private final AtomicLong pinnedBytes = new AtomicLong(); + private long fineBytes; + private boolean closed; + private final String segment; + private final int segmentMaxDoc; + + SegmentIVFVectorsReader(SegmentReadState state) throws IOException { + String name = segment = state.segmentInfo.name, sfx = state.segmentSuffix; + segmentMaxDoc = state.segmentInfo.maxDoc(); + byte[] id = state.segmentInfo.getId(); + String metaName = IndexFileNames.segmentFileName(name, sfx, META_EXTENSION); + try (ChecksumIndexInput meta = state.directory.openChecksumInput(metaName)) { + Throwable prior = null; + try { + checkIndexHeader(meta, META_CODEC_NAME, VERSION_CURRENT, VERSION_CURRENT, id, sfx); + for (int number = meta.readInt(); number != -1; number = meta.readInt()) { + FieldInfo info = state.fieldInfos.fieldInfo(number); + if (info == null) throw new CorruptIndexException("invalid field number " + number, meta); + fields.put(info.name, new Field(meta, info)); + } + } catch (Throwable t) { + prior = t; + } finally { + CodecUtil.checkFooter(meta, prior); + } + } + String dataName = IndexFileNames.segmentFileName(name, sfx, DATA_EXTENSION); + data = state.directory.openInput(dataName, state.context); + try { + checkIndexHeader(data, DATA_CODEC_NAME, VERSION_CURRENT, VERSION_CURRENT, id, sfx); + CodecUtil.retrieveChecksum(data); + } catch (Throwable t) { + IOUtils.closeWhileSuppressingExceptions(t, data); + throw t; + } + uring = openUring(state, dataName); + for (Field field : fields.values()) fineBytes += field.sections[2] - field.sections[1]; + Uring.addFine(fineBytes); + if (state.context.context() != IOContext.Context.MERGE) { + for (Field field : fields.values()) field.maybePin(); + } + } + + /** Opens batched fine-record reads for a data file of a file-system directory, if possible. */ + private static Uring openUring(SegmentReadState state, String name) { + if (FilterDirectory.unwrap(state.directory) instanceof FSDirectory fs) { + try { + return Uring.open(fs.getDirectory().resolve(name)); + } catch (IOException _) { + // fine records are read through the mapped input instead + } + } + return null; + } + + final class Field { + final VectorSimilarityFunction similarity; + final FineTier fineTier; + final int dim, nlist, count, nprobe, spillBits, recordLen, coarseBytes, docIdOffset; + final long rotationSeed; + final long[] sections; + final DirectMonotonicReader.Meta postingOffsets; + final HadamardRotation rotation; + final FineCodec fine; + + RandomAccessInput records, slotDocs; + MemorySegment recordsSeg; // the whole mapped records section, or null + MemorySegmentAccessInput coarseAccess; + MemorySegment coarseSeg; + // Off-heap copies published by the pinning thread; searches use the mapped sections until then. + volatile MemorySegment pinnedCoarse, pinnedSlotDocs; + private final AtomicBoolean pinning = new AtomicBoolean(); + private volatile long lastPinAttempt; + int[] cellStart, ordToSlot, ordToDoc, allCells, primaryCells; + float[][] centroids; + private volatile CentroidCodes codes; + CentroidGraph graph; + + Field(ChecksumIndexInput meta, FieldInfo info) throws IOException { + similarity = info.getVectorSimilarityFunction(); + fineTier = FineTier.values()[meta.readByte()]; + dim = meta.readVInt(); + nlist = meta.readVInt(); + count = meta.readVInt(); + rotationSeed = meta.readLong(); + nprobe = meta.readVInt(); + spillBits = meta.readVInt(); + // Sections: centroids, records, coarse, graph, ordToSlot, slotDoc, posting offsets. + sections = new long[7]; + for (int s = 0; s < sections.length; s++) sections[s] = meta.readVLong(); + postingOffsets = nlist == 0 ? null : loadMeta(meta, nlist + 1, DIRECT_MONOTONIC_BLOCK_SHIFT); + if (dim != info.getVectorDimension()) throw new CorruptIndexException("dimension", meta); + rotation = HadamardRotation.create(dim, rotationSeed); + fine = new FineCodec(fineTier, dim); + recordLen = CodeRecord.length(fine.codeBytes); + coarseBytes = Nitrox2.bytesPerVector(dim); + docIdOffset = fine.codeBytes; + } + + private RandomAccessInput section(int s) throws IOException { + return data.randomAccessSlice(sections[s], sections[s + 1] - sections[s]); + } + + synchronized Field open() throws IOException { + if (records != null) return this; + if (section(2) instanceof MemorySegmentAccessInput in) { + coarseAccess = in; + coarseSeg = segmentOrNull(in, 0, in.length()); + } + cellStart = new int[nlist + 1]; + if (nlist > 0) { + long postings = sections[sections.length - 1]; + RandomAccessInput tail = data.randomAccessSlice(postings, data.length() - postings); + var offsets = DirectMonotonicReader.getInstance(postingOffsets, tail); + for (int c = 0; c <= nlist; c++) cellStart[c] = (int) (offsets.get(c) / Integer.BYTES); + } + slotDocs = section(5); + IndexInput all = data.clone(); + centroids = new float[nlist][dim]; + all.seek(sections[0]); + for (float[] centroid : centroids) all.readFloats(centroid, 0, dim); + records = section(1); + if (records instanceof MemorySegmentAccessInput in) { + recordsSeg = segmentOrNull(in, 0, in.length()); + } + return this; + } + + /** + * Queues the off-heap copy of the coarse and slot-to-document sections when the pin budget has + * room, retrying at most once per second: a merged segment is pinned once its sources close. + * Searches never wait for the copy. + */ + void maybePin() { + if (count == 0 || pinnedCoarse != null || pinning.get()) return; + long now = System.nanoTime(); + if (lastPinAttempt != 0 && now - lastPinAttempt < 1_000_000_000L) return; + if (pinning.getAndSet(true)) return; + lastPinAttempt = now == 0 ? 1 : now; + long coarseLength = sections[3] - sections[2]; + long docsLength = sections[6] - sections[5]; + if (Uring.reservePinned(coarseLength + docsLength) == false) { + pinning.set(false); + return; + } + pinnedBytes.addAndGet(coarseLength + docsLength); + Uring.PINNER.execute( + () -> { + try { + if (docsLength > 0) pinnedSlotDocs = pin(sections[5], docsLength); + pinnedCoarse = pin(sections[2], coarseLength); + } catch (IOException | RuntimeException _) { + // the reader closed mid-copy; its reservation is released on close + } + }); + } + + /** Copies a section off-heap, where the page cache cannot evict it. */ + private MemorySegment pin(long offset, long length) throws IOException { + MemorySegment copy = pinned.allocate(length, 64); // cache-line aligned for the SIMD scans + IndexInput in = data.clone(); + in.seek(offset); + byte[] chunk = new byte[1 << 20]; + for (long at = 0; at < length; at += chunk.length) { + int n = (int) Math.min(chunk.length, length - at); + in.readBytes(chunk, 0, n); + MemorySegment.copy(chunk, 0, copy, ValueLayout.JAVA_BYTE, at, n); + } + return copy; + } + + private int docAt(RandomAccessInput docs, int slot) throws IOException { + MemorySegment pinned = pinnedSlotDocs; + return pinned != null ? pinned.getAtIndex(INT_LE, slot) : docs.readInt((long) slot * 4); + } + + private synchronized void loadOrdToSlot() throws IOException { + if (ordToSlot != null) return; + int[] slots = new int[count]; + IndexInput all = data.clone(); + all.seek(sections[4]); + all.readInts(slots, 0, count); + ordToSlot = slots; + } + + private synchronized void loadOrdinalMappings() throws IOException { + if (ordToDoc != null) return; + loadOrdToSlot(); + int[] docs = new int[count]; + for (int ord = 0; ord < count; ord++) docs[ord] = docAt(slotDocs, ordToSlot[ord]); + ordToDoc = docs; + } + + /** + * Returns a mapped slice, or null when no single mapping covers it. The slice is rebased as a + * plain native segment scoped to this reader, so the SIMD kernels see one segment type whether + * or not a section is pinned: once they have seen both, every coarse scan runs ~16% slower. + */ + @SuppressWarnings("restricted") + private MemorySegment segmentOrNull(MemorySegmentAccessInput in, long offset, long length) { + if (length == 0) return null; + try { + MemorySegment mapped = in.segmentSliceOrNull(offset, length); + if (mapped == null) return null; + return MemorySegment.ofAddress(mapped.address()).reinterpret(length, pinned, null); + } catch (IOException _) { + return null; + } + } + + /** + * Returns the mapping holding one cell run: the whole coarse section when mappable, otherwise a + * slice rebased to the run, or null when neither is mapped. Offsets start at {@code runBase}. + */ + private MemorySegment coarseRun(int slotBase, int rows) { + if (coarseSeg != null || coarseAccess == null) return coarseSeg; + return segmentOrNull(coarseAccess, (long) slotBase * coarseBytes, (long) rows * coarseBytes); + } + + private long runBase(int slotBase) { + return coarseSeg == null ? 0 : (long) slotBase * coarseBytes; + } + + private synchronized void loadCodes() throws IOException { + if (codes != null) return; + allCells = new int[nlist]; + for (int c = 0; c < nlist; c++) allCells[c] = c; + long graphLength = sections[4] - sections[3]; + graph = graphLength == 0 ? null : CentroidGraph.read(section(3), dim, graphLength); + codes = new CentroidCodes(centroids, dim, fine); + } + + synchronized int cellOf(int ord) throws IOException { + loadOrdToSlot(); + if (primaryCells == null) { + primaryCells = new int[count]; + for (int o = 0; o < count; o++) { + primaryCells[o] = records.readInt((long) ordToSlot[o] * recordLen + docIdOffset + 4); + } + } + return primaryCells[ord]; + } + + /** + * Per-query state for cell selection, coarse admission, deduplication, and fine reranking. Each + * search opens its own slices, which shadow the field's: positional reads through a shared + * slice are not thread-safe on every directory implementation. + */ + final class Search { + final RandomAccessInput records = section(1), coarse = section(2); + final RandomAccessInput slotDocs = section(5); + // One snapshot per query, so a copy published mid-query cannot mix offsets. + final MemorySegment pinnedCoarse = Field.this.pinnedCoarse; + final float[] rotated; + final byte[] qCode; + final FineCodec.Query fine; + final KnnCollector collector; + final Scratch scratch = Scratch.LOCAL.get(); + final int bins = coarseBytes * 8 + 2, shortlist, pool; + final int[] histogram = scratch.histogram = ArrayUtil.growNoCopy(scratch.histogram, bins); + int size, admitted, threshold = bins - 1; + Candidates gathered; + + Search(float[] target, KnnCollector collector, int k) throws IOException { + maybePin(); + this.collector = collector; + shortlist = (int) Math.min(count, rerankCount(k)); + pool = Math.multiplyExact(shortlist, 1 + spillBits); + scratch.reserve(shortlist); + Arrays.fill(histogram, 0, bins, 0); + if (codes == null) loadCodes(); + rotated = new float[dim]; + qCode = new byte[coarseBytes]; + rotation.rotate(VectorUtil.l2normalize(ArrayUtil.copyOfSubArray(target, 0, dim)), rotated); + Nitrox2.encode(rotated, dim, qCode, 0); + fine = Field.this.fine.query(rotated, similarity); + } + + /** A rerank of {@code from}'s shortlist on this thread, reusing its encoded query. */ + private Search(Search from, KnnCollector collector) throws IOException { + this.collector = collector; + shortlist = from.shortlist; + pool = from.pool; + rotated = from.rotated; + qCode = from.qCode; + fine = from.fine; + } + + /** Fine-reranks {@code slots}, taken from this search's shortlist, into {@code collector}. */ + void rerankInto(int[] slots, KnnCollector collector) throws IOException { + new Search(this, collector).rerank(slots, slots.length); + } + + void run(AcceptDocs acceptDocs) throws IOException { + var strategy = collector.getSearchStrategy(); + if (strategy instanceof SearchStrategy s) run(s.numProbes, s.probeMargin, acceptDocs); + else run(nprobe, DEFAULT_PROBE_MARGIN, acceptDocs); + } + + private void run(int probe, float margin, AcceptDocs acceptDocs) throws IOException { + probe = Math.min(probe, nlist); + Bits accept = acceptDocs == null ? null : acceptDocs.bits(); + if (accept instanceof BitSet filter) { + int cost = acceptDocs.cost(); + double parity = + Math.sqrt((double) shortlist * cellStart[nlist] * coarseBytes / recordLen); + if (cost > (int) Math.max(shortlist, Math.min(Integer.MAX_VALUE, parity))) { + filteredScan(selectCells(probe, 1f), filter, cost, probe); + return; + } + boolean dense = count == segmentMaxDoc; + if (dense) loadOrdToSlot(); + else loadOrdinalMappings(); + int[] slots = new int[64]; + int n = 0; + DocIdSetIterator accepted = acceptDocs.iterator(); + for (int doc = accepted.nextDoc(); doc != NO_MORE_DOCS; doc = accepted.nextDoc()) { + int ord = dense ? doc : Arrays.binarySearch(ordToDoc, doc); + if (ord < 0) continue; + slots = ArrayUtil.grow(slots, n + 1); + slots[n++] = ordToSlot[ord]; + } + if (collector == null) admitAll(slots, n); + else rerank(slots, n); + } else { + // A segment no larger than the rerank pool scans every slot: its mostly empty cells would + // spend probes on cells holding nothing. + scan(cellStart[nlist] <= pool ? allCells : selectCells(probe, margin), accept); + } + } + + private int[] selectCells(int probe, float margin) { + int[] candidates = allCells; + int got = nlist; + if (graph != null) { + candidates = scratch.candidates = ArrayUtil.growNoCopy(scratch.candidates, nlist); + int[] coarse = scratch.coarse = ArrayUtil.growNoCopy(scratch.coarse, nlist); + int ef = Math.max(CentroidGraph.MIN_EF, probe * CentroidGraph.EF_MULTIPLIER); + got = graph.search(qCode, ef, candidates, coarse); + int cap = Math.max(VERIFY_MIN, probe * VERIFY_MULTIPLIER); + if (cap < got) { + int[] counts = new int[bins]; + for (int i = 0; i < got; i++) counts[coarse[i]]++; + int below = 0, bound = 0; + while (below + counts[bound] <= cap) below += counts[bound++]; + int n = 0, ties = cap - below; + for (int i = 0; i < got && n < cap; i++) { + if (coarse[i] < bound || (coarse[i] == bound && ties-- > 0)) { + candidates[n] = candidates[i]; + coarse[n++] = coarse[i]; + } + } + got = n; + } + } + long[] ranked = rank(candidates, got); Review Comment: Is this another place where we don't actually need a full rank, but rather a partially ordered top-k list? If so, maybe another candidate for a quick-select style algo? -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
