Skip to content
2 changes: 2 additions & 0 deletions lucene/CHANGES.txt
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,8 @@ New Features

* GITHUB#16383: Add fp16 vector encoding support. (Pulkit Gupta)

* GITHUB#16473: Add scalar quantization support in Fp16 vector encoding. (Pulkit Gupta)

Improvements
---------------------
* GITHUB#15704: Replace LinkedList with more efficient data structure. (Renato Haeberli)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,43 @@ public RandomVectorScorer getRandomVectorScorer(
public RandomVectorScorer getRandomVectorScorer(
VectorSimilarityFunction similarityFunction, KnnVectorValues vectorValues, short[] target)
throws IOException {
if (vectorValues instanceof QuantizedByteVectorValues qv) {
FlatVectorsScorer.checkDimensions(target.length, qv.dimension());
OptimizedScalarQuantizer quantizer = qv.getQuantizer();
ScalarEncoding scalarEncoding = qv.getScalarEncoding();
byte[] scratch = new byte[scalarEncoding.getDiscreteDimensions(qv.dimension())];
final byte[] targetQuantized;
if (scalarEncoding.isAsymmetric() == false) {
targetQuantized = scratch;
} else {
// This is asymmetric quantization, we will pack the vector
targetQuantized = new byte[scalarEncoding.getQueryPackedLength(scratch.length)];
}
// Inflate the fp16 query to fp32 and normalize there; quantization operates on fp32.
float[] copy = new float[target.length];
for (int i = 0; i < target.length; i++) {
copy[i] = Float.float16ToFloat(target[i]);
}
if (similarityFunction == COSINE) {
VectorUtil.l2normalize(copy);
}
var targetCorrectiveTerms =
quantizer.scalarQuantize(copy, scratch, scalarEncoding.getQueryBits(), qv.getCentroid());
// for asymmetric encodings with 4-bit query, we need to transpose the nibbles for fast
// scoring comparisons
if (scalarEncoding == ScalarEncoding.SINGLE_BIT_QUERY_NIBBLE
|| scalarEncoding == ScalarEncoding.DIBIT_QUERY_NIBBLE) {
OptimizedScalarQuantizer.transposeHalfByte(scratch, targetQuantized);
}
return new RandomVectorScorer.AbstractRandomVectorScorer(qv) {
@Override
public float score(int node) throws IOException {
return quantizedScore(
targetQuantized, targetCorrectiveTerms, qv, node, similarityFunction);
}
};
}
// It is possible to get to this branch during initial indexing and flush
return nonQuantizedDelegate.getRandomVectorScorer(similarityFunction, vectorValues, target);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -224,7 +224,26 @@ public RandomVectorScorer getRandomVectorScorer(String field, byte[] target) thr

@Override
public RandomVectorScorer getRandomVectorScorer(String field, short[] target) throws IOException {
return rawVectorsReader.getRandomVectorScorer(field, target);
FieldEntry fi = fields.get(field);
if (fi == null) {
return null;
}
return vectorScorer.getRandomVectorScorer(
fi.similarityFunction,
OffHeapScalarQuantizedVectorValues.load(
fi.ordToDocDISIReaderConfiguration,
fi.dimension,
fi.size,
new OptimizedScalarQuantizer(fi.similarityFunction),
fi.scalarEncoding,
fi.similarityFunction,
vectorScorer,
fi.centroid,
fi.centroidDP,
fi.vectorDataOffset,
fi.vectorDataLength,
quantizedVectorData),
target);
}

@Override
Expand Down Expand Up @@ -289,7 +308,57 @@ public ByteVectorValues getByteVectorValues(String field) throws IOException {

@Override
public Float16VectorValues getFloat16VectorValues(String field) throws IOException {
return rawVectorsReader.getFloat16VectorValues(field);
FieldEntry fi = fields.get(field);
if (fi == null) {
return null;
}
if (fi.vectorEncoding != VectorEncoding.FLOAT16) {
throw new IllegalArgumentException(
"field=\""
+ field
+ "\" is encoded as: "
+ fi.vectorEncoding
+ " expected: "
+ VectorEncoding.FLOAT16);
}

Float16VectorValues rawFloat16VectorValues = rawVectorsReader.getFloat16VectorValues(field);

OffHeapScalarQuantizedVectorValues sqvv =
OffHeapScalarQuantizedVectorValues.load(
fi.ordToDocDISIReaderConfiguration,
fi.dimension,
fi.size,
new OptimizedScalarQuantizer(fi.similarityFunction),
fi.scalarEncoding,
fi.similarityFunction,
vectorScorer,
fi.centroid,
fi.centroidDP,
fi.vectorDataOffset,
fi.vectorDataLength,
quantizedVectorData);

if (rawFloat16VectorValues.size() == 0) {
// The raw float16 vectors were dropped, so reads reconstruct values by dequantizing. Pair
// that view with sqvv so scorer() scores in quantized space while vectorValue() and
// rescorer() dequantize.
Float16VectorValues dequantizedRawVectorValues =
OffHeapScalarQuantizedFloat16VectorValues.load(
fi.ordToDocDISIReaderConfiguration,
fi.dimension,
fi.size,
fi.scalarEncoding,
fi.similarityFunction,
vectorScorer,
fi.centroid,
fi.vectorDataOffset,
fi.vectorDataLength,
quantizedVectorData);
return new ScalarQuantizedFloat16VectorValues(dequantizedRawVectorValues, sqvv);
}

return new ScalarQuantizedFloat16VectorValues(rawFloat16VectorValues, sqvv);
}

@Override
Expand All @@ -302,11 +371,25 @@ public void search(String field, byte[] target, KnnCollector knnCollector, Accep
public void search(String field, float[] target, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
if (knnCollector.k() == 0) return;
final RandomVectorScorer scorer = getRandomVectorScorer(field, target);
exhaustiveBulkScore(getRandomVectorScorer(field, target), knnCollector, acceptDocs);
}

@Override
public void search(String field, short[] target, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
if (knnCollector.k() == 0) return;
exhaustiveBulkScore(getRandomVectorScorer(field, target), knnCollector, acceptDocs);
}

/**
* Scores every accepted vector with the given scorer, collecting into {@code knnCollector}.
* Scoring happens in batches of {@link #EXHAUSTIVE_BULK_SCORE_ORDS} ordinals.
*/
private static void exhaustiveBulkScore(
RandomVectorScorer scorer, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
if (scorer == null) return;
Bits acceptedOrds = scorer.getAcceptOrds(acceptDocs.bits());
// if k is larger than the number of vectors we expect to visit in an HNSW search,
// we can just iterate over all vectors and collect them.
int[] ords = new int[EXHAUSTIVE_BULK_SCORE_ORDS];
float[] scores = new float[EXHAUSTIVE_BULK_SCORE_ORDS];
int numOrds = 0;
Expand Down Expand Up @@ -339,12 +422,6 @@ public void search(String field, float[] target, KnnCollector knnCollector, Acce
}
}

@Override
public void search(String field, short[] target, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
rawVectorsReader.search(field, target, knnCollector, acceptDocs);
}

@Override
public void close() throws IOException {
IOUtils.close(quantizedVectorData, rawVectorsReader);
Expand All @@ -366,10 +443,7 @@ public Map<String, Long> getOffHeapByteSize(FieldInfo fieldInfo) {
var raw = rawVectorsReader.getOffHeapByteSize(fieldInfo);
var fieldEntry = fields.get(fieldInfo.name);
if (fieldEntry == null) {
// Only FLOAT32 fields are scalar-quantized by this format; BYTE and FLOAT16 fields are
// stored raw by the delegate and therefore have no quantized field entry here.
assert fieldInfo.getVectorEncoding() == VectorEncoding.BYTE
|| fieldInfo.getVectorEncoding() == VectorEncoding.FLOAT16;
assert fieldInfo.getVectorEncoding() == VectorEncoding.BYTE;
return raw;
}
var quant = Map.of(VECTOR_DATA_EXTENSION, fieldEntry.vectorDataLength());
Expand Down Expand Up @@ -442,14 +516,16 @@ public QuantizedByteVectorValues getQuantizedVectorValues(String field) throws I
if (fi == null) {
return null;
}
if (fi.vectorEncoding != VectorEncoding.FLOAT32) {
if (fi.vectorEncoding.isFloatingPoint() == false) {
throw new IllegalArgumentException(
"field=\""
+ field
+ "\" is encoded as: "
+ fi.vectorEncoding
+ " expected: "
+ VectorEncoding.FLOAT32);
+ VectorEncoding.FLOAT32
+ " or "
+ VectorEncoding.FLOAT16);
}
return OffHeapScalarQuantizedVectorValues.load(
fi.ordToDocDISIReaderConfiguration,
Expand Down Expand Up @@ -692,4 +768,66 @@ QuantizedByteVectorValues getQuantizedVectorValues() throws IOException {
return quantizedVectorValues;
}
}

/** Vector values holding raw and quantized vector values */
protected static final class ScalarQuantizedFloat16VectorValues extends Float16VectorValues {
private final Float16VectorValues rawVectorValues;
private final QuantizedByteVectorValues quantizedVectorValues;

ScalarQuantizedFloat16VectorValues(
Float16VectorValues rawVectorValues, QuantizedByteVectorValues quantizedVectorValues) {
this.rawVectorValues = rawVectorValues;
this.quantizedVectorValues = quantizedVectorValues;
}

@Override
public int dimension() {
return rawVectorValues.dimension();
}

@Override
public int size() {
return rawVectorValues.size();
}

@Override
public short[] vectorValue(int ord) throws IOException {
return rawVectorValues.vectorValue(ord);
}

@Override
public ScalarQuantizedFloat16VectorValues copy() throws IOException {
return new ScalarQuantizedFloat16VectorValues(
rawVectorValues.copy(), quantizedVectorValues.copy());
}

@Override
public Bits getAcceptOrds(Bits acceptDocs) {
return rawVectorValues.getAcceptOrds(acceptDocs);
}

@Override
public int ordToDoc(int ord) {
return rawVectorValues.ordToDoc(ord);
}

@Override
public DocIndexIterator iterator() {
return rawVectorValues.iterator();
}

@Override
public VectorScorer scorer(short[] query) throws IOException {
return quantizedVectorValues.scorer(query);
}

@Override
public VectorScorer rescorer(short[] target) throws IOException {
return rawVectorValues.rescorer(target);
}

QuantizedByteVectorValues getQuantizedVectorValues() throws IOException {
return quantizedVectorValues;
}
}
}
Loading
Loading