From c7395b8c1f008d09e8639e58e8e7384821818f71 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Thu, 16 Jul 2026 10:18:24 +0000 Subject: [PATCH 01/13] Add random-access int8 vector value support Extract shared random-access behavior into the generic VectorValues interface while preserving RandomAccessVectorValues for float vectors. Add RandomAccessByteVectorValues and its list-backed implementation to support native ByteSequence vectors without float32 conversion. Update BuildScoreProvider to use the generalized thread-local supplier. --- .../ListRandomAccessByteVectorValues.java | 70 +++++++++++++++++ .../graph/RandomAccessByteVectorValues.java | 37 +++++++++ .../graph/RandomAccessVectorValues.java | 47 ++--------- .../jbellis/jvector/graph/VectorValues.java | 77 +++++++++++++++++++ .../graph/similarity/BuildScoreProvider.java | 7 +- 5 files changed, 196 insertions(+), 42 deletions(-) create mode 100644 jvector-base/src/main/java/io/github/jbellis/jvector/graph/ListRandomAccessByteVectorValues.java create mode 100644 jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessByteVectorValues.java create mode 100644 jvector-base/src/main/java/io/github/jbellis/jvector/graph/VectorValues.java diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/ListRandomAccessByteVectorValues.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/ListRandomAccessByteVectorValues.java new file mode 100644 index 000000000..134b13546 --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/ListRandomAccessByteVectorValues.java @@ -0,0 +1,70 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.graph; + +import io.github.jbellis.jvector.vector.types.ByteSequence; + +import java.util.List; + +/** + * A List-backed implementation of the {@link RandomAccessByteVectorValues} interface. + *

+ * It is acceptable to provide this class to a GraphBuilder, and then continue + * to add vectors to the backing List as you add to the graph. + *

+ * This will be as threadsafe as the provided List. + */ +public class ListRandomAccessByteVectorValues implements RandomAccessByteVectorValues { + private final List> vectors; + private final int dimension; + + /** + * Construct a new instance of {@link ListRandomAccessByteVectorValues}. + * + * @param vectors a (potentially mutable) list of byte vectors. + * @param dimension the dimension of the vectors. + */ + public ListRandomAccessByteVectorValues(List> vectors, int dimension) { + this.vectors = vectors; + this.dimension = dimension; + } + + @Override + public int size() { + return vectors.size(); + } + + @Override + public int dimension() { + return dimension; + } + + @Override + public ByteSequence getVector(int nodeId) { + return vectors.get(nodeId); + } + + @Override + public boolean isValueShared() { + return false; + } + + @Override + public ListRandomAccessByteVectorValues copy() { + return this; + } +} diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessByteVectorValues.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessByteVectorValues.java new file mode 100644 index 000000000..fa81dd705 --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessByteVectorValues.java @@ -0,0 +1,37 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.graph; + +import io.github.jbellis.jvector.vector.types.ByteSequence; + +/** + * Provides random access to byte (int8) vectors by dense ordinal. + *

+ * This is the byte-vector parallel to {@link RandomAccessVectorValues}. + * Both extend the common super-interface {@link VectorValues}. + * It is used by graph-based index builders and searchers that operate natively + * on int8 vectors without a float32 round-trip. + */ +public interface RandomAccessByteVectorValues extends VectorValues> { + + /** + * Creates a new copy of this {@link RandomAccessByteVectorValues}. + * Un-shared implementations may simply return {@code this}. + */ + @Override + RandomAccessByteVectorValues copy(); +} diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessVectorValues.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessVectorValues.java index eb8f6df24..3e36a6f82 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessVectorValues.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RandomAccessVectorValues.java @@ -30,14 +30,15 @@ import io.github.jbellis.jvector.vector.types.VectorFloat; import java.util.function.Supplier; -import java.util.logging.Logger; /** - * Provides random access to vectors by dense ordinal. This interface is used by graph-based + * Provides random access to float32 vectors by dense ordinal. This interface is used by graph-based * implementations of KNN search. + *

+ * For int8 vectors see {@link RandomAccessByteVectorValues}. + * Both extend the common super-interface {@link VectorValues}. */ -public interface RandomAccessVectorValues { - Logger LOG = Logger.getLogger(RandomAccessVectorValues.class.getName()); +public interface RandomAccessVectorValues extends VectorValues> { /** * Return the number of vector values. @@ -46,23 +47,9 @@ public interface RandomAccessVectorValues { * (1) implementing a threadsafe, un-shared RAVV, where `copy` returns `this`, or * (2) implementing a fixed-size RAVV. */ + @Override int size(); - /** Return the dimension of the returned vector values */ - int dimension(); - - /** - * Return the vector value indexed at the given ordinal. - * - *

For performance, implementations are free to re-use the same object across invocations. - * That is, you will get back the same VectorFloat<?> - * reference (for instance) for every requested ordinal. If you want to use those values across - * calls, you should make a copy. - * - * @param nodeId a valid ordinal, ≥ 0 and < {@link #size()}. - */ - VectorFloat getVector(int nodeId); - @Deprecated default VectorFloat vectorValue(int targetOrd) { return getVector(targetOrd); @@ -78,12 +65,6 @@ default void getVectorInto(int node, VectorFloat destinationVector, int offse destinationVector.copyFrom(getVector(node), 0, offset, dimension()); } - /** - * @return true iff the vector returned by `getVector` is shared. A shared vector will - * only be valid until the next call to getVector overwrites it. - */ - boolean isValueShared(); - /** * Creates a new copy of this {@link RandomAccessVectorValues}. This is helpful when you need to * access different values at once, to avoid overwriting the underlying float vector returned by @@ -91,23 +72,9 @@ default void getVectorInto(int node, VectorFloat destinationVector, int offse *

* Un-shared implementations may simply return `this`. */ + @Override RandomAccessVectorValues copy(); - /** - * Returns a supplier of thread-local copies of the RAVV. - */ - default Supplier threadLocalSupplier() { - if (!isValueShared()) { - return () -> this; - } - - if (this instanceof AutoCloseable) { - LOG.warning("RAVV is shared and implements AutoCloseable; threadLocalSupplier() may lead to leaks"); - } - var tl = ExplicitThreadLocal.withInitial(this::copy); - return tl::get; - } - /** * Convenience method to create an ExactScoreFunction for reranking. The resulting function is NOT thread-safe. */ diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/VectorValues.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/VectorValues.java new file mode 100644 index 000000000..f490aac23 --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/VectorValues.java @@ -0,0 +1,77 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.graph; + +import io.github.jbellis.jvector.util.ExplicitThreadLocal; + +import java.util.function.Supplier; +import java.util.logging.Logger; + +/** + * Common super-interface for random access to vectors by dense ordinal. + *

+ * {@code V} is the vector element type — {@code VectorFloat} for float32 vectors + * (see {@link RandomAccessVectorValues}) and {@code ByteSequence} for int8 vectors + * (see {@link RandomAccessByteVectorValues}). + */ +public interface VectorValues { + Logger LOG = Logger.getLogger(VectorValues.class.getName()); + + /** Return the number of vector values. */ + int size(); + + /** Return the dimension of the returned vector values. */ + int dimension(); + + /** + * Return the vector value indexed at the given ordinal. + *

+ * For performance, implementations are free to re-use the same object across invocations. + * If you need to retain the value across calls, make a copy. + * + * @param nodeId a valid ordinal, ≥ 0 and < {@link #size()}. + */ + V getVector(int nodeId); + + /** + * @return true iff the vector returned by {@link #getVector} is shared across calls. + * A shared vector is only valid until the next call to {@link #getVector} overwrites it. + */ + boolean isValueShared(); + + /** + * Creates a new copy of this instance. + * Un-shared implementations may simply return {@code this}. + */ + VectorValues copy(); + + /** + * Returns a supplier of thread-local copies of this instance. + */ + @SuppressWarnings("unchecked") + default Supplier> threadLocalSupplier() { + if (!isValueShared()) { + return () -> this; + } + + if (this instanceof AutoCloseable) { + LOG.warning("VectorValues is shared and implements AutoCloseable; threadLocalSupplier() may lead to leaks"); + } + var tl = ExplicitThreadLocal.withInitial(this::copy); + return tl::get; + } +} diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java index 1049069de..987bfba52 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java @@ -25,6 +25,7 @@ import io.github.jbellis.jvector.vector.VectorizationProvider; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import java.util.function.Supplier; /** * Encapsulates comparing node distances for GraphIndexBuilder. @@ -106,8 +107,10 @@ static BuildScoreProvider randomAccessScoreProvider(RandomAccessVectorValues rav static BuildScoreProvider randomAccessScoreProvider(RandomAccessVectorValues ravv, VectorSimilarityFunction similarityFunction) { // We need two sources of vectors in order to perform diversity check comparisons without // colliding. ThreadLocalSupplier makes this a no-op if the RAVV is actually un-shared. - var vectors = ravv.threadLocalSupplier(); - var vectorsCopy = ravv.threadLocalSupplier(); + var vectorsRaw = ravv.threadLocalSupplier(); + var vectorsCopyRaw = ravv.threadLocalSupplier(); + Supplier vectors = () -> (RandomAccessVectorValues) vectorsRaw.get(); + Supplier vectorsCopy = () -> (RandomAccessVectorValues) vectorsCopyRaw.get(); return new BuildScoreProvider() { @Override From ad86e042801d777f90ae741a7372d9dbb0ff0384 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Thu, 16 Jul 2026 10:18:41 +0000 Subject: [PATCH 02/13] Add byte-similarity methods to VectorUtilSupport / VectorUtil / DefaultVectorUtilSupport MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add three fundamental signed int8 vector similarity operations to the vectorization support layer so they participate in the same provider dispatch as float similarity. Bytes are treated as signed int8 (Java byte range -128..127). VectorUtilSupport — three new abstract methods: float dotProduct(ByteSequence a, ByteSequence b) float squareDistance(ByteSequence a, ByteSequence b) float cosine(ByteSequence a, ByteSequence b) VectorUtil — three new public static delegates: dotProduct(ByteSequence, ByteSequence) -> impl.dotProduct squareL2Distance(ByteSequence, ByteSequence) -> impl.squareDistance cosine(ByteSequence, ByteSequence) -> impl.cosine DefaultVectorUtilSupport — scalar loop implementations: dotProduct: accumulate (int)a.get(i) * (int)b.get(i), return as float. squareDistance: accumulate (diff * diff) for each signed byte difference. cosine: dot / sqrt(normA * normB) using per-element float promotion. PanamaVectorUtilSupport — scalar stub overrides identical to Default, so jvector-twenty compiles without requiring a SIMD implementation now. SIMD optimisation of byte similarity is a future concern. --- .../vector/DefaultVectorUtilSupport.java | 31 +++++++++++++++++++ .../jbellis/jvector/vector/VectorUtil.java | 15 +++++++++ .../jvector/vector/VectorUtilSupport.java | 9 ++++++ .../vector/PanamaVectorUtilSupport.java | 31 +++++++++++++++++++ 4 files changed, 86 insertions(+) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/DefaultVectorUtilSupport.java b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/DefaultVectorUtilSupport.java index 5843dc5f6..e5f7b0824 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/DefaultVectorUtilSupport.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/DefaultVectorUtilSupport.java @@ -338,6 +338,37 @@ public float assembleAndSumPQ( return res; } + @Override + public float dotProduct(ByteSequence a, ByteSequence b) { + float sum = 0; + for (int i = 0; i < a.length(); i++) { + sum += (int) a.get(i) * (int) b.get(i); + } + return sum; + } + + @Override + public float squareDistance(ByteSequence a, ByteSequence b) { + float sum = 0; + for (int i = 0; i < a.length(); i++) { + float diff = a.get(i) - b.get(i); + sum += diff * diff; + } + return sum; + } + + @Override + public float cosine(ByteSequence a, ByteSequence b) { + float dot = 0, normA = 0, normB = 0; + for (int i = 0; i < a.length(); i++) { + float ai = a.get(i), bi = b.get(i); + dot += ai * bi; + normA += ai * ai; + normB += bi * bi; + } + return (float) (dot / Math.sqrt(normA * normB)); + } + @Override public int hammingDistance(long[] v1, long[] v2) { int hd = 0; diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtil.java b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtil.java index 744d5ec75..01550f264 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtil.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtil.java @@ -174,6 +174,21 @@ public static float assembleAndSumPQ(VectorFloat data, int subspaceCount, Byt return impl.assembleAndSumPQ(data, subspaceCount, dataOffsets1, dataOffsetsOffset1, dataOffsets2, dataOffsetsOffset2, clusterCount); } + /** Returns the dot product of two signed int8 byte vectors. */ + public static float dotProduct(ByteSequence a, ByteSequence b) { + return impl.dotProduct(a, b); + } + + /** Returns the sum of squared differences of two signed int8 byte vectors. */ + public static float squareL2Distance(ByteSequence a, ByteSequence b) { + return impl.squareDistance(a, b); + } + + /** Returns the cosine similarity of two signed int8 byte vectors. */ + public static float cosine(ByteSequence a, ByteSequence b) { + return impl.cosine(a, b); + } + public static int hammingDistance(long[] v1, long[] v2) { return impl.hammingDistance(v1, v2); } diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtilSupport.java b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtilSupport.java index 118f16ca6..01a706405 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtilSupport.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/VectorUtilSupport.java @@ -130,6 +130,15 @@ public interface VectorUtilSupport { */ float assembleAndSumPQ(VectorFloat codebookPartialSums, int subspaceCount, ByteSequence vector1Ordinals, int vector1OrdinalOffset, ByteSequence node2Ordinals, int node2OrdinalOffset, int clusterCount); + /** Calculates the dot product of two signed int8 byte vectors. */ + float dotProduct(ByteSequence a, ByteSequence b); + + /** Returns the sum of squared differences of two signed int8 byte vectors. */ + float squareDistance(ByteSequence a, ByteSequence b); + + /** Returns the cosine similarity of two signed int8 byte vectors. */ + float cosine(ByteSequence a, ByteSequence b); + int hammingDistance(long[] v1, long[] v2); void calculatePartialSums(VectorFloat codebook, int codebookIndex, int size, int clusterCount, VectorFloat query, int offset, VectorSimilarityFunction vsf, VectorFloat partialSums); diff --git a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java index 22e0d2c60..f95d4bf6b 100644 --- a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java +++ b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java @@ -984,6 +984,37 @@ float assembleAndSumPQ_512( * @param b The second array * @return The Hamming distance */ + @Override + public float dotProduct(ByteSequence a, ByteSequence b) { + float sum = 0; + for (int i = 0; i < a.length(); i++) { + sum += (int) a.get(i) * (int) b.get(i); + } + return sum; + } + + @Override + public float squareDistance(ByteSequence a, ByteSequence b) { + float sum = 0; + for (int i = 0; i < a.length(); i++) { + float diff = a.get(i) - b.get(i); + sum += diff * diff; + } + return sum; + } + + @Override + public float cosine(ByteSequence a, ByteSequence b) { + float dot = 0, normA = 0, normB = 0; + for (int i = 0; i < a.length(); i++) { + float ai = a.get(i), bi = b.get(i); + dot += ai * bi; + normA += ai * ai; + normB += bi * bi; + } + return (float) (dot / Math.sqrt(normA * normB)); + } + @Override public int hammingDistance(long[] a, long[] b) { var sum = LongVector.zero(LongVector.SPECIES_PREFERRED); From ee69adfae4f1c40d84c2f5dc511c277e582d4b4d Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Thu, 16 Jul 2026 10:19:03 +0000 Subject: [PATCH 03/13] Add ByteVectorSimilarityFunction enum New enum in jvector-base/.../vector/ parallel to VectorSimilarityFunction but operating on ByteSequence, delegating to the VectorUtil byte methods from Sub-Task 2. Three variants with return values normalised to [0,1] matching VectorSimilarityFunction conventions (higher = more similar): EUCLIDEAN: 1 / (1 + squaredL2 / (n * 255^2)) Normalises by the maximum possible squared distance between two signed int8 vectors (255^2 per dimension) so the result stays in (0,1] regardless of dimension. DOT_PRODUCT: (1 + dot / (n * 127^2)) / 2 Normalises by the maximum possible dot product magnitude (127^2 per dimension) before applying the (1+x)/2 mapping so the result stays in [0,1] regardless of dimension or whether vectors are unit-norm. For already unit-norm int8 vectors (e.g. Cohere, OpenAI reduced-precision) prefer COSINE. COSINE: (1 + cosine(v1, v2)) / 2 Cosine is inherently bounded to [-1,1] so no extra normalisation is needed. --- .../vector/ByteVectorSimilarityFunction.java | 76 +++++++++++++++++++ 1 file changed, 76 insertions(+) create mode 100644 jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java new file mode 100644 index 000000000..2390343ea --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java @@ -0,0 +1,76 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.vector; + +import io.github.jbellis.jvector.vector.types.ByteSequence; + +/** + * Vector similarity function for signed int8 (byte) vectors; parallel to + * {@link VectorSimilarityFunction} but operating on {@link ByteSequence}. + *

+ * Bytes are treated as signed int8 values (Java's {@code byte} is already signed, range −128..127). + * Return-value conventions match {@link VectorSimilarityFunction}: higher is more similar. + */ +public enum ByteVectorSimilarityFunction { + + /** + * Euclidean similarity normalised to {@code (0, 1]}. + * Raw squared L2 is divided by {@code n * 255^2} (the maximum possible squared distance + * between two signed int8 vectors) before the {@code 1 / (1 + x)} mapping, so the result + * is always in (0, 1] regardless of dimension. + */ + EUCLIDEAN { + @Override + public float compare(ByteSequence v1, ByteSequence v2) { + float maxSquaredDist = v1.length() * (255.0f * 255.0f); + return 1.0f / (1.0f + VectorUtil.squareL2Distance(v1, v2) / maxSquaredDist); + } + }, + + /** + * Dot product normalised to {@code [0, 1]}. + * Raw int8 dot product is divided by {@code n * 127^2} (the maximum possible magnitude) + * before applying the {@code (1 + x) / 2} mapping, so the result is always in [0, 1] + * regardless of dimension or whether the vectors are unit-norm. + * For already unit-norm int8 vectors (e.g. Cohere, OpenAI reduced-precision) prefer {@link #COSINE}. + */ + DOT_PRODUCT { + @Override + public float compare(ByteSequence v1, ByteSequence v2) { + float maxMagnitude = v1.length() * (127.0f * 127.0f); + return (1.0f + VectorUtil.dotProduct(v1, v2) / maxMagnitude) / 2.0f; + } + }, + + /** Cosine similarity: {@code (1 + cosine(v1, v2)) / 2} */ + COSINE { + @Override + public float compare(ByteSequence v1, ByteSequence v2) { + return (1.0f + VectorUtil.cosine(v1, v2)) / 2.0f; + } + }; + + /** + * Calculates a similarity score between the two int8 vectors. + * Higher values correspond to closer vectors. + * + * @param v1 a byte vector + * @param v2 another byte vector, of the same dimension + * @return the similarity score + */ + public abstract float compare(ByteSequence v1, ByteSequence v2); +} From f4a48e4bfe3893390615b97c2473cf65f868fd96 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Mon, 7 Sep 2026 08:58:12 +0000 Subject: [PATCH 04/13] Add BuildScoreProvider.byteVectorScoreProvider and searchProviderFor(ByteSequence) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add a new static factory byteVectorScoreProvider(RandomAccessByteVectorValues, ByteVectorSimilarityFunction) that performs exact byte×byte scoring with no float32 round-trip. The returned BuildScoreProvider implements isExact()=true, approximateCentroid(), searchProviderFor(ByteSequence), searchProviderFor(int), diversityProviderFor(int), and diversityScoreFunctionFor(int), using two independent threadLocalSupplier() handles for thread-safe concurrent builds. Also adds a default searchProviderFor(ByteSequence) to the interface that throws UnsupportedOperationException, so existing float-based providers are unaffected. --- .../graph/similarity/BuildScoreProvider.java | 79 +++++++++++++++++++ 1 file changed, 79 insertions(+) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java index 987bfba52..634a6f8cf 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java @@ -16,13 +16,16 @@ package io.github.jbellis.jvector.graph.similarity; +import io.github.jbellis.jvector.graph.RandomAccessByteVectorValues; import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import io.github.jbellis.jvector.graph.RemappedRandomAccessVectorValues; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.quantization.BQVectors; import io.github.jbellis.jvector.quantization.PQVectors; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorUtil; import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; import java.util.function.Supplier; @@ -60,6 +63,20 @@ public interface BuildScoreProvider { */ SearchScoreProvider searchProviderFor(VectorFloat vector); + /** + * Create a search score provider to use *internally* during construction, for a byte (int8) query vector. + *

+ * The default implementation throws {@link UnsupportedOperationException}; only + * {@link #byteVectorScoreProvider} overrides it. + * + * @param vector the int8 query vector to provide similarity scores against + */ + default SearchScoreProvider searchProviderFor(ByteSequence vector) { + throw new UnsupportedOperationException( + "This BuildScoreProvider does not support byte-vector queries; " + + "use byteVectorScoreProvider() to construct a byte-vector builder"); + } + /** * Create a search score provider to use *internally* during construction. *

@@ -214,6 +231,68 @@ public VectorFloat approximateCentroid() { }; } + /** + * Returns a BSP that performs exact score comparisons using the given + * {@link RandomAccessByteVectorValues} and {@link ByteVectorSimilarityFunction}. + * All scoring is byte×byte with no float32 round-trip. + */ + static BuildScoreProvider byteVectorScoreProvider(RandomAccessByteVectorValues ravv, ByteVectorSimilarityFunction bvsf) { + var vectors = ravv.threadLocalSupplier(); + var vectorsCopy = ravv.threadLocalSupplier(); + + return new BuildScoreProvider() { + @Override + public boolean isExact() { + return true; + } + + @Override + public VectorFloat approximateCentroid() { + var vv = vectors.get(); + var centroid = vts.createFloatVector(vv.dimension()); + for (int i = 0; i < vv.size(); i++) { + var v = vv.getVector(i); + for (int d = 0; d < vv.dimension(); d++) { + centroid.set(d, centroid.get(d) + v.get(d)); + } + } + VectorUtil.scale(centroid, 1.0f / vv.size()); + return centroid; + } + + @Override + public SearchScoreProvider searchProviderFor(VectorFloat vector) { + throw new UnsupportedOperationException( + "byteVectorScoreProvider does not support float query vectors; use searchProviderFor(int node)"); + } + + @Override + public SearchScoreProvider searchProviderFor(ByteSequence vector) { + var vc = vectorsCopy.get(); + var sf = (ScoreFunction.ExactScoreFunction) node2 -> bvsf.compare(vector, vc.getVector(node2)); + return new DefaultSearchScoreProvider(sf); + } + + @Override + public SearchScoreProvider searchProviderFor(int node1) { + var v = vectors.get().getVector(node1); + return searchProviderFor(v); + } + + @Override + public SearchScoreProvider diversityProviderFor(int node1) { + return searchProviderFor(node1); + } + + @Override + public ScoreFunction diversityScoreFunctionFor(int node1) { + var v = vectors.get().getVector(node1); + var vc = vectorsCopy.get(); + return (ScoreFunction.ExactScoreFunction) node2 -> bvsf.compare(v, vc.getVector(node2)); + } + }; + } + static BuildScoreProvider bqBuildScoreProvider(BQVectors bqv) { return new BuildScoreProvider() { @Override From a2db47be0a31dee847c2ef7efe8d4ea3ec611058 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Mon, 7 Sep 2026 08:58:29 +0000 Subject: [PATCH 05/13] Wire byteVectorScoreProvider into GraphIndexBuilder Add builder(RandomAccessByteVectorValues, ByteVectorSimilarityFunction, int M) and builder(..., List maxDegrees) factory overloads so callers can build a graph over int8 byte vectors without touching the builder core. Generalize build(RandomAccessVectorValues) to build(VectorValues) and switch the parallel addGraphNode loop to use scoreProvider.searchProviderFor(node) directly, removing the float-only getVector path that was incompatible with byte-vector score providers. Add addGraphNode(int node, ByteSequence vector) as a public single-node insertion entry point for byte vectors. Update TestVectorGraph to add explicit casts (RandomAccessVectorValues, VectorSimilarityFunction) on the null arguments so the compiler resolves the correct overload now that the new byte-vector builder() overloads exist. --- .../jvector/graph/GraphIndexBuilder.java | 58 +++++++++++++++++-- .../jvector/graph/TestVectorGraph.java | 2 +- 2 files changed, 55 insertions(+), 5 deletions(-) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/GraphIndexBuilder.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/GraphIndexBuilder.java index c6f65910c..d1cf15213 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/GraphIndexBuilder.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/GraphIndexBuilder.java @@ -23,6 +23,7 @@ import io.github.jbellis.jvector.graph.SearchResult.NodeScore; import io.github.jbellis.jvector.graph.diversity.VamanaDiversityProvider; import io.github.jbellis.jvector.graph.similarity.BuildScoreProvider; +import io.github.jbellis.jvector.graph.VectorValues; import io.github.jbellis.jvector.graph.similarity.ScoreFunction; import io.github.jbellis.jvector.graph.similarity.SearchScoreProvider; import io.github.jbellis.jvector.management.CompressionType; @@ -32,7 +33,9 @@ import io.github.jbellis.jvector.quantization.PQVectors; import io.github.jbellis.jvector.quantization.ProductQuantization; import io.github.jbellis.jvector.util.*; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -114,6 +117,11 @@ private static BuildScoreProvider getBuildScoreProvider(RandomAccessVectorValues } } + private static BuildScoreProvider getBuildScoreProvider(RandomAccessByteVectorValues vectorValues, ByteVectorSimilarityFunction similarityFunction) { + // Byte vectors do not support PQ/BQ build-time compression; always use the direct provider. + return BuildScoreProvider.byteVectorScoreProvider(vectorValues, similarityFunction); + } + /** * Reads all the vectors from vector values, builds a graph connecting them by their dense * ordinals, using the given hyperparameter settings, and returns the resulting graph. @@ -407,6 +415,31 @@ public static Builder builder(RandomAccessVectorValues vectorValues, VectorSimil return builder(vectorValues, similarityFunction, List.of(M)); } + /** + * Entry point for the fluent builder for int8 byte vectors, resolving the {@link BuildScoreProvider} + * from raw byte vectors via {@link BuildScoreProvider#byteVectorScoreProvider}. + * + * @param vectorValues the int8 byte vectors whose relations are represented by the graph + * @param similarityFunction the byte-vector similarity function to score vectors with + * @param maxDegrees the maximum number of connections a node can have in each layer; if fewer entries + * are specified than the number of layers, the last entry is used for all remaining layers. + */ + public static Builder builder(RandomAccessByteVectorValues vectorValues, ByteVectorSimilarityFunction similarityFunction, List maxDegrees) { + return new Builder(getBuildScoreProvider(vectorValues, similarityFunction), vectorValues.dimension(), maxDegrees); + } + + /** + * Entry point for the fluent builder for int8 byte vectors, for the common case of a single + * (non-hierarchical) max degree. + * + * @param vectorValues the int8 byte vectors whose relations are represented by the graph + * @param similarityFunction the byte-vector similarity function to score vectors with + * @param M the maximum number of connections a node can have + */ + public static Builder builder(RandomAccessByteVectorValues vectorValues, ByteVectorSimilarityFunction similarityFunction, int M) { + return builder(vectorValues, similarityFunction, List.of(M)); + } + /** * Entry point for the fluent builder, building from an existing {@link MutableGraphIndex} (e.g. one just * loaded from disk) rather than constructing a fresh {@link OnHeapGraphIndex}. {@code addHierarchy} is not @@ -695,13 +728,17 @@ public static GraphIndexBuilder rescore(GraphIndexBuilder other, BuildScoreProvi return newBuilder; } - public ImmutableGraphIndex build(RandomAccessVectorValues ravv) { - var vv = ravv.threadLocalSupplier(); + /** + * Builds the graph from any {@link VectorValues} source — works for both + * {@link RandomAccessVectorValues} (float32) and {@link RandomAccessByteVectorValues} (int8). + * The score provider supplied at construction time determines how vectors are compared. + */ + public ImmutableGraphIndex build(VectorValues ravv) { int size = ravv.size(); simdExecutor.submit(() -> { IntStream.range(0, size).parallel().forEach(node -> { - addGraphNode(node, vv.get().getVector(node)); + addGraphNode(node, scoreProvider.searchProviderFor(node)); }); }).join(); @@ -852,6 +889,19 @@ public long addGraphNode(int node, VectorFloat vector) { return addGraphNode(node, ssp); } + /** + * Inserts a node with the given int8 byte vector into the graph. + * + * @param node the node ID to add + * @param vector the byte vector to add + * @return an estimate of the number of extra bytes used by the graph after adding the given node + * @throws UnsupportedOperationException if this builder was not constructed with a byte-vector score provider + */ + public long addGraphNode(int node, ByteSequence vector) { + var ssp = scoreProvider.searchProviderFor(vector); + return addGraphNode(node, ssp); + } + /** * Inserts a node with the given vector value to the graph. * @@ -1337,4 +1387,4 @@ public static ImmutableGraphIndex buildAndMergeNewNodes(RandomAccessReader in, return builder.getGraph(); } } -} \ No newline at end of file +} diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java index 52bdc872a..b784cd7d4 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java @@ -433,7 +433,7 @@ public void testGraphIndexBuilderInvalid() { public void testGraphIndexBuilderInvalid(boolean addHierarchy) { assertThrows(NullPointerException.class, - () -> new GraphIndexBuilder(null, null, 0, 0, 1.0f, 1.0f, addHierarchy)); + () -> new GraphIndexBuilder((RandomAccessVectorValues) null, (VectorSimilarityFunction) null, 0, 0, 1.0f, 1.0f, addHierarchy)); // M must be > 0 assertThrows(IllegalArgumentException.class, () -> { From 15fef54e482cf8e48a3925bef76582a563ec6701 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Mon, 20 Jul 2026 07:19:21 +0000 Subject: [PATCH 06/13] Vectorize ByteSequence similarity metrics with SIMD Replace scalar loop implementations of dotProduct, squareDistance, and cosine for ByteSequence with Panama Vector API implementations in PanamaVectorUtilSupport, and native AVX-512/AVX2 kernels wired through NativeVectorUtilSupport -> NativeSimdOps JNI bindings. Java (Panama) path: - Widen signed bytes to int32 via B2I conversion, accumulate products in IntVector lanes, then reduce. - Dispatch on PREFERRED_BIT_SIZE: 512-bit: load 16 bytes (SPECIES_128) -> IntVector.SPECIES_512 256-bit: load 8 bytes (SPECIES_64) -> IntVector.SPECIES_256 128-bit: scalar fallback - cosine variants accumulate dot/norm products in long after reduction to avoid int32 overflow on large vectors. Native path (C++): - New kernels in jvector_simd_kernels.cpp and jvector_avx3_dl_kernels.cpp: dot_product_i8, euclidean_i8, cosine_i8 using Highway SIMD (AVX-512 / AVX2 dispatch). - Registered in jvector_simd_kernel_list.h and exported via jvector_simd.cpp. - Microbenchmarks added in bench_similarity_i8.cpp. - C++ unit tests added in test_similarity_i8.cpp using a prime-length (107-element) vector to exercise tail handling. Java tests: - TestVectorizationProvider.testSimilarityMetricsByte cross-checks SIMD results against scalar DefaultVectorUtilSupport baseline. --- .../vector/NativeVectorUtilSupport.java | 24 ++ .../jvector/vector/cnative/NativeSimdOps.java | 186 +++++++++ .../native/benchmarks/bench_similarity_i8.cpp | 117 ++++++ jvector-native/src/main/native/meson.build | 6 +- .../native/src/jvector_avx3_dl_kernels.cpp | 353 +++++++++++++++++- .../src/main/native/src/jvector_simd.cpp | 9 +- .../native/src/jvector_simd_kernel_list.h | 6 +- .../main/native/src/jvector_simd_kernels.cpp | 169 +++++++++ .../src/main/native/tests/test_helpers.cpp | 22 ++ .../src/main/native/tests/test_helpers.h | 6 +- .../main/native/tests/test_similarity_i8.cpp | 229 ++++++++++++ .../vector/TestVectorizationProvider.java | 34 ++ .../vector/PanamaVectorUtilSupport.java | 269 +++++++++++-- 13 files changed, 1397 insertions(+), 33 deletions(-) create mode 100644 jvector-native/src/main/native/benchmarks/bench_similarity_i8.cpp create mode 100644 jvector-native/src/main/native/tests/test_similarity_i8.cpp diff --git a/jvector-native/src/main/java/io/github/jbellis/jvector/vector/NativeVectorUtilSupport.java b/jvector-native/src/main/java/io/github/jbellis/jvector/vector/NativeVectorUtilSupport.java index 4b627a244..f20ba7805 100644 --- a/jvector-native/src/main/java/io/github/jbellis/jvector/vector/NativeVectorUtilSupport.java +++ b/jvector-native/src/main/java/io/github/jbellis/jvector/vector/NativeVectorUtilSupport.java @@ -61,6 +61,30 @@ public String getMaxIsaEnv() { return ptr.reinterpret(Long.MAX_VALUE).getString(0); } + @Override + public float dotProduct(ByteSequence a, ByteSequence b) { + return NativeSimdOps.dot_product_i8( + ((MemorySegmentByteSequence) a).get(), (long) a.offset(), + ((MemorySegmentByteSequence) b).get(), (long) b.offset(), + (long) a.length()); + } + + @Override + public float squareDistance(ByteSequence a, ByteSequence b) { + return NativeSimdOps.euclidean_i8( + ((MemorySegmentByteSequence) a).get(), (long) a.offset(), + ((MemorySegmentByteSequence) b).get(), (long) b.offset(), + (long) a.length()); + } + + @Override + public float cosine(ByteSequence a, ByteSequence b) { + return NativeSimdOps.cosine_i8( + ((MemorySegmentByteSequence) a).get(), (long) a.offset(), + ((MemorySegmentByteSequence) b).get(), (long) b.offset(), + (long) a.length()); + } + @Override protected FloatVector fromVectorFloat(VectorSpecies SPEC, VectorFloat vector, int offset) { return FloatVector.fromMemorySegment(SPEC, ((MemorySegmentVectorFloat) vector).get(), vector.offset(offset), ByteOrder.LITTLE_ENDIAN); diff --git a/jvector-native/src/main/java/io/github/jbellis/jvector/vector/cnative/NativeSimdOps.java b/jvector-native/src/main/java/io/github/jbellis/jvector/vector/cnative/NativeSimdOps.java index d822468a6..19fe50a72 100644 --- a/jvector-native/src/main/java/io/github/jbellis/jvector/vector/cnative/NativeSimdOps.java +++ b/jvector-native/src/main/java/io/github/jbellis/jvector/vector/cnative/NativeSimdOps.java @@ -2506,6 +2506,192 @@ public static void nvq_shuffle_query_in_place_8bit(MemorySegment vector, long le } } + private static class dot_product_i8 { + public static final FunctionDescriptor DESC = FunctionDescriptor.of( + NativeSimdOps.C_FLOAT, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_LONG + ); + + public static final MemorySegment ADDR = NativeSimdOps.findOrThrow("dot_product_i8"); + + public static final MethodHandle HANDLE = Linker.nativeLinker().downcallHandle(ADDR, DESC, Linker.Option.critical(true)); + } + + /** + * Function descriptor for: + * {@snippet lang=c : + * float dot_product_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static FunctionDescriptor dot_product_i8$descriptor() { + return dot_product_i8.DESC; + } + + /** + * Downcall method handle for: + * {@snippet lang=c : + * float dot_product_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MethodHandle dot_product_i8$handle() { + return dot_product_i8.HANDLE; + } + + /** + * Address for: + * {@snippet lang=c : + * float dot_product_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MemorySegment dot_product_i8$address() { + return dot_product_i8.ADDR; + } + + /** + * {@snippet lang=c : + * float dot_product_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static float dot_product_i8(MemorySegment a, long aoffset, MemorySegment b, long boffset, long length) { + var mh$ = dot_product_i8.HANDLE; + try { + if (TRACE_DOWNCALLS) { + traceDowncall("dot_product_i8", a, aoffset, b, boffset, length); + } + return (float)mh$.invokeExact(a, aoffset, b, boffset, length); + } catch (Throwable ex$) { + throw new AssertionError("should not reach here", ex$); + } + } + + private static class euclidean_i8 { + public static final FunctionDescriptor DESC = FunctionDescriptor.of( + NativeSimdOps.C_FLOAT, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_LONG + ); + + public static final MemorySegment ADDR = NativeSimdOps.findOrThrow("euclidean_i8"); + + public static final MethodHandle HANDLE = Linker.nativeLinker().downcallHandle(ADDR, DESC, Linker.Option.critical(true)); + } + + /** + * Function descriptor for: + * {@snippet lang=c : + * float euclidean_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static FunctionDescriptor euclidean_i8$descriptor() { + return euclidean_i8.DESC; + } + + /** + * Downcall method handle for: + * {@snippet lang=c : + * float euclidean_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MethodHandle euclidean_i8$handle() { + return euclidean_i8.HANDLE; + } + + /** + * Address for: + * {@snippet lang=c : + * float euclidean_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MemorySegment euclidean_i8$address() { + return euclidean_i8.ADDR; + } + + /** + * {@snippet lang=c : + * float euclidean_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static float euclidean_i8(MemorySegment a, long aoffset, MemorySegment b, long boffset, long length) { + var mh$ = euclidean_i8.HANDLE; + try { + if (TRACE_DOWNCALLS) { + traceDowncall("euclidean_i8", a, aoffset, b, boffset, length); + } + return (float)mh$.invokeExact(a, aoffset, b, boffset, length); + } catch (Throwable ex$) { + throw new AssertionError("should not reach here", ex$); + } + } + + private static class cosine_i8 { + public static final FunctionDescriptor DESC = FunctionDescriptor.of( + NativeSimdOps.C_FLOAT, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_POINTER, + NativeSimdOps.C_LONG, + NativeSimdOps.C_LONG + ); + + public static final MemorySegment ADDR = NativeSimdOps.findOrThrow("cosine_i8"); + + public static final MethodHandle HANDLE = Linker.nativeLinker().downcallHandle(ADDR, DESC, Linker.Option.critical(true)); + } + + /** + * Function descriptor for: + * {@snippet lang=c : + * float cosine_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static FunctionDescriptor cosine_i8$descriptor() { + return cosine_i8.DESC; + } + + /** + * Downcall method handle for: + * {@snippet lang=c : + * float cosine_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MethodHandle cosine_i8$handle() { + return cosine_i8.HANDLE; + } + + /** + * Address for: + * {@snippet lang=c : + * float cosine_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static MemorySegment cosine_i8$address() { + return cosine_i8.ADDR; + } + + /** + * {@snippet lang=c : + * float cosine_i8(const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length) + * } + */ + public static float cosine_i8(MemorySegment a, long aoffset, MemorySegment b, long boffset, long length) { + var mh$ = cosine_i8.HANDLE; + try { + if (TRACE_DOWNCALLS) { + traceDowncall("cosine_i8", a, aoffset, b, boffset, length); + } + return (float)mh$.invokeExact(a, aoffset, b, boffset, length); + } catch (Throwable ex$) { + throw new AssertionError("should not reach here", ex$); + } + } + private static class jvector_simd_get_active_isa { public static final FunctionDescriptor DESC = FunctionDescriptor.of( NativeSimdOps.C_POINTER ); diff --git a/jvector-native/src/main/native/benchmarks/bench_similarity_i8.cpp b/jvector-native/src/main/native/benchmarks/bench_similarity_i8.cpp new file mode 100644 index 000000000..fd9af793a --- /dev/null +++ b/jvector-native/src/main/native/benchmarks/bench_similarity_i8.cpp @@ -0,0 +1,117 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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. + */ + +// Google Benchmark micro-benchmarks for the int8 vector similarity kernels: +// dot_product_i8, euclidean_i8, cosine_i8 +// +// Parameterised over the realistic embedding dimensions used in production: +// 128, 256, 512, 1024, 1536, 3072 +// +// Build (requires google-benchmark installed or available via pkg-config): +// meson setup build && ninja -C build bench_simd_kernels +// +// Run: +// ./build/bench_simd_kernels [--benchmark_filter=] + +#include +#include +#include + +#include "jvector_simd.h" + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +// Deterministic, non-zero int8 vector: values cycle through a signed range to +// avoid degenerate all-zero inputs while staying within [-128, 127]. +static std::vector make_i8_vec(size_t n, int8_t seed) +{ + std::vector v(n); + for (size_t i = 0; i < n; ++i) { + int val = seed + static_cast(i % 127); + if (i % 3 == 0) val = -val; + // clamp to [-127, 127] to keep vectors non-degenerate for cosine + if (val > 127) val = 127; + if (val < -127) val = -127; + v[i] = static_cast(val); + } + return v; +} + +// Benchmark sizes matching production embedding dimensions. +static const std::vector kBenchSizes = {128, 256, 512, 1024, 1536, 3072}; + +// --------------------------------------------------------------------------- +// dot_product_i8 +// --------------------------------------------------------------------------- + +static void BM_dot_product_i8(benchmark::State& state) +{ + const size_t n = static_cast(state.range(0)); + auto a = make_i8_vec(n, 7); + auto b = make_i8_vec(n, 13); + + for (auto _ : state) { + float result = dot_product_i8(a.data(), 0, b.data(), 0, n); + benchmark::DoNotOptimize(result); + } + + state.SetItemsProcessed(state.iterations() * static_cast(n)); + state.SetBytesProcessed(state.iterations() * static_cast(n) * 2 * sizeof(int8_t)); +} +BENCHMARK(BM_dot_product_i8)->ArgsProduct({kBenchSizes}); + +// --------------------------------------------------------------------------- +// euclidean_i8 +// --------------------------------------------------------------------------- + +static void BM_euclidean_i8(benchmark::State& state) +{ + const size_t n = static_cast(state.range(0)); + auto a = make_i8_vec(n, 7); + auto b = make_i8_vec(n, 13); + + for (auto _ : state) { + float result = euclidean_i8(a.data(), 0, b.data(), 0, n); + benchmark::DoNotOptimize(result); + } + + state.SetItemsProcessed(state.iterations() * static_cast(n)); + state.SetBytesProcessed(state.iterations() * static_cast(n) * 2 * sizeof(int8_t)); +} +BENCHMARK(BM_euclidean_i8)->ArgsProduct({kBenchSizes}); + +// --------------------------------------------------------------------------- +// cosine_i8 +// --------------------------------------------------------------------------- + +static void BM_cosine_i8(benchmark::State& state) +{ + const size_t n = static_cast(state.range(0)); + auto a = make_i8_vec(n, 7); + auto b = make_i8_vec(n, 13); + + for (auto _ : state) { + float result = cosine_i8(a.data(), 0, b.data(), 0, n); + benchmark::DoNotOptimize(result); + } + + state.SetItemsProcessed(state.iterations() * static_cast(n)); + state.SetBytesProcessed(state.iterations() * static_cast(n) * 2 * sizeof(int8_t)); +} +BENCHMARK(BM_cosine_i8)->ArgsProduct({kBenchSizes}); + diff --git a/jvector-native/src/main/native/meson.build b/jvector-native/src/main/native/meson.build index 42fec1ada..d992716f3 100644 --- a/jvector-native/src/main/native/meson.build +++ b/jvector-native/src/main/native/meson.build @@ -131,6 +131,7 @@ if gtest_dep.found() sources : [ 'tests/test_helpers.cpp', 'tests/test_similarity.cpp', + 'tests/test_similarity_i8.cpp', 'tests/test_elementwise.cpp', 'tests/test_cpu_features.cpp', ], @@ -154,7 +155,10 @@ gbench_dep = dependency('benchmark', required: false) if gbench_dep.found() executable( 'bench_simd_kernels', - sources : 'benchmarks/bench_similarity_f32.cpp', + sources : [ + 'benchmarks/bench_similarity_f32.cpp', + 'benchmarks/bench_similarity_i8.cpp', + ], dependencies: [vectorutil_dep, gbench_dep], cpp_args : ['-O3'], ) diff --git a/jvector-native/src/main/native/src/jvector_avx3_dl_kernels.cpp b/jvector-native/src/main/native/src/jvector_avx3_dl_kernels.cpp index ae8ab73bb..11e7ed33d 100644 --- a/jvector-native/src/main/native/src/jvector_avx3_dl_kernels.cpp +++ b/jvector-native/src/main/native/src/jvector_avx3_dl_kernels.cpp @@ -25,13 +25,358 @@ // VNNI, VBMI, VBMI2, IFMA, BITALG, VPOPCNTDQ, GFNI, VAES, VPCLMULQDQ // // Compiled with -march=icelake-server. -// Highway will select HWY_AVX3_DL as the static target. +// +// This file uses raw Intel AVX-512 intrinsics directly — NO Google Highway — +// so we get exactly the instructions we intend with zero abstraction overhead. + +#include +#include +#include +#include // AVX-512 + VNNI intrinsics #include "jvector_simd.h" -#include "hwy/highway.h" -#include "assert_hwy_targets.h" -namespace hn = hwy::HWY_NAMESPACE; +// ============================================================================= +// Register naming convention +// zmm = 512-bit (16 × int32, 32 × int16, 64 × int8) +// ymm = 256-bit (32 × int8, 16 × int16) +// xmm = 128-bit (16 × int8, 8 × int16) +// +// VNNI instructions used +// ───────────────────────────────────────────────────────────────────────────── +// VPDPBUSD zmm_acc, zmm_a, zmm_b +// For each group of 4 adjacent lanes (i×4 .. i×4+3): +// acc[i] += (u8)a[i×4+0] * (i8)b[i×4+0] +// + (u8)a[i×4+1] * (i8)b[i×4+1] +// + (u8)a[i×4+2] * (i8)b[i×4+2] +// + (u8)a[i×4+3] * (i8)b[i×4+3] +// → 16 i32 accumulations, 64 int8 products per zmm register per cycle. +// Latency: 3 cycles. Throughput: 1/cycle (two ports on Ice Lake). +// +// VPDPWSSD zmm_acc, zmm_a, zmm_b +// For each group of 2 adjacent i16 lanes (i×2, i×2+1): +// acc[i] += (i16)a[i×2+0] * (i16)b[i×2+0] +// + (i16)a[i×2+1] * (i16)b[i×2+1] +// → 16 i32 accumulations, 32 int16 products per zmm per cycle. +// Latency: 3 cycles. Throughput: 1/cycle. +// +// Signed i8 × signed i8 using VPDPBUSD +// ───────────────────────────────────────────────────────────────────────────── +// VPDPBUSD requires operand A to be unsigned. For signed inputs we apply the +// standard bias trick: +// (a + 128) is always non-negative, so we use it as the unsigned operand. +// (a+128) * b = a*b + 128*b → a*b = VPDPBUSD(a+128, b) - 128 * sum(b) +// +// The bias (128*sum(b)) is constant per zmm load of b, computed as: +// _mm512_dpwssd_epi32(zero, b, set1_epi16(128)) [reuse VPDPWSSD] +// and subtracted once per iteration from the accumulator. +// +// This adds one VPDPWSSD + one VPADDD per iteration, which is negligible +// compared to the main VPDPBUSD throughput. +// +// Unrolling strategy +// ───────────────────────────────────────────────────────────────────────────── +// With 3-cycle VPDPBUSD latency and 1/cycle throughput (ports 0+5), we need +// at least 4 independent accumulator chains to keep the ports saturated: +// issued cycle 0: port 0 ← acc0 +// issued cycle 1: port 5 ← acc1 +// issued cycle 2: port 0 ← acc2 +// issued cycle 3: port 5 ← acc3 (acc0 writeback done, cycle 3) +// 4× unrolling fully hides the 3-cycle latency. +// ============================================================================= namespace AVX3_DL { +// --------------------------------------------------------------------------- +// Horizontal reduce: sum all 16 int32 lanes of a zmm register. +// _mm512_reduce_add_epi32 emits the optimal fold-down sequence; the compiler +// schedules it across surrounding instructions better than manual shuffles. +// --------------------------------------------------------------------------- +static inline int32_t hsum_epi32(__m512i v) +{ + return _mm512_reduce_add_epi32(v); +} + +// --------------------------------------------------------------------------- +// dot_product_i8 — VPDPBUSD with bias correction for signed i8 × signed i8 +// --------------------------------------------------------------------------- +// +// Algorithm +// acc = VPDPBUSD(acc, a_u8, b_i8) where a_u8 = a + 128 +// bias = VPDPWSSD(bias, b_i8, 128) accumulates 128 * sum(b) +// result = hsum(acc) - hsum(bias) +// +// 4× unrolled (256 bytes/iteration) to saturate both ICX VNNI ports and +// fully hide the 3-cycle VPDPBUSD latency. +float dot_product_i8(const int8_t * __restrict__ a, size_t aoffset, + const int8_t * __restrict__ b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + + const __m512i bias128 = _mm512_set1_epi16(128); + + __m512i acc0 = _mm512_setzero_si512(), acc1 = _mm512_setzero_si512(); + __m512i acc2 = _mm512_setzero_si512(), acc3 = _mm512_setzero_si512(); + __m512i bias0 = _mm512_setzero_si512(), bias1 = _mm512_setzero_si512(); + __m512i bias2 = _mm512_setzero_si512(), bias3 = _mm512_setzero_si512(); + + size_t i = 0; + for (; i + 256 <= length; i += 256) { + __m512i va0 = _mm512_loadu_si512(a + i + 0); + __m512i va1 = _mm512_loadu_si512(a + i + 64); + __m512i va2 = _mm512_loadu_si512(a + i + 128); + __m512i va3 = _mm512_loadu_si512(a + i + 192); + __m512i vb0 = _mm512_loadu_si512(b + i + 0); + __m512i vb1 = _mm512_loadu_si512(b + i + 64); + __m512i vb2 = _mm512_loadu_si512(b + i + 128); + __m512i vb3 = _mm512_loadu_si512(b + i + 192); + + // Flip sign bit: maps signed [-128,127] → unsigned [0,255]. + const __m512i flip = _mm512_set1_epi8(-128); + __m512i au0 = _mm512_add_epi8(va0, flip); + __m512i au1 = _mm512_add_epi8(va1, flip); + __m512i au2 = _mm512_add_epi8(va2, flip); + __m512i au3 = _mm512_add_epi8(va3, flip); + + // VPDPBUSD: acc[i] += (u8)au[4i+k] * (i8)vb[4i+k], k=0..3 + acc0 = _mm512_dpbusd_epi32(acc0, au0, vb0); + acc1 = _mm512_dpbusd_epi32(acc1, au1, vb1); + acc2 = _mm512_dpbusd_epi32(acc2, au2, vb2); + acc3 = _mm512_dpbusd_epi32(acc3, au3, vb3); + + // Bias: promote vb to i16 then compute 128 * sum(vb) using VPDPWSSD. + // Each 64-byte zmm of int8 is split into two 512-bit i16 vectors. + __m512i vb0_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 0))); + __m512i vb0_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 32))); + __m512i vb1_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 64))); + __m512i vb1_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 96))); + __m512i vb2_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 128))); + __m512i vb2_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 160))); + __m512i vb3_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 192))); + __m512i vb3_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 224))); + + bias0 = _mm512_dpwssd_epi32(bias0, vb0_lo, bias128); + bias0 = _mm512_dpwssd_epi32(bias0, vb0_hi, bias128); + bias1 = _mm512_dpwssd_epi32(bias1, vb1_lo, bias128); + bias1 = _mm512_dpwssd_epi32(bias1, vb1_hi, bias128); + bias2 = _mm512_dpwssd_epi32(bias2, vb2_lo, bias128); + bias2 = _mm512_dpwssd_epi32(bias2, vb2_hi, bias128); + bias3 = _mm512_dpwssd_epi32(bias3, vb3_lo, bias128); + bias3 = _mm512_dpwssd_epi32(bias3, vb3_hi, bias128); + } + __m512i acc = _mm512_add_epi32(_mm512_add_epi32(acc0, acc1), + _mm512_add_epi32(acc2, acc3)); + __m512i bias = _mm512_add_epi32(_mm512_add_epi32(bias0, bias1), + _mm512_add_epi32(bias2, bias3)); + + // Single-zmm tail (residual 64-byte blocks). + for (; i + 64 <= length; i += 64) { + __m512i va = _mm512_loadu_si512(a + i); + __m512i vb = _mm512_loadu_si512(b + i); + __m512i au = _mm512_add_epi8(va, _mm512_set1_epi8(-128)); + acc = _mm512_dpbusd_epi32(acc, au, vb); + + // Promote vb to i16 for bias calculation + __m512i vb_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i))); + __m512i vb_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 32))); + bias = _mm512_dpwssd_epi32(bias, vb_lo, bias128); + bias = _mm512_dpwssd_epi32(bias, vb_hi, bias128); + } + + int32_t result = hsum_epi32(acc) - hsum_epi32(bias); + + // Scalar tail. + for (; i < length; i++) + result += (int32_t)a[i] * (int32_t)b[i]; + + return (float)result; +} + +// --------------------------------------------------------------------------- +// euclidean_i8 — VPMOVSXBW sign-extend + VPDPWSSD squared differences +// --------------------------------------------------------------------------- +// +// Each 64-byte zmm block is processed as two 32-byte halves: +// _mm512_cvtepi8_epi16(__m256i) = VPMOVSXBW: sign-extends 32×i8 → 32×i16 +// diff = da - db (i16 subtraction, no overflow since range is [-255,255]) +// acc = VPDPWSSD(acc, diff, diff) +// +// 4× unrolled (256 bytes/iteration). +float euclidean_i8(const int8_t * __restrict__ a, size_t aoffset, + const int8_t * __restrict__ b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + + __m512i acc0 = _mm512_setzero_si512(), acc1 = _mm512_setzero_si512(); + __m512i acc2 = _mm512_setzero_si512(), acc3 = _mm512_setzero_si512(); + + size_t i = 0; + for (; i + 256 <= length; i += 256) { +#define EUCL_BLOCK(off, acc_var) \ + { \ + const int8_t *ap = a + i + (off), *bp = b + i + (off); \ + __m512i da_lo = _mm512_cvtepi8_epi16( \ + _mm256_loadu_si256(reinterpret_cast(ap))); \ + __m512i db_lo = _mm512_cvtepi8_epi16( \ + _mm256_loadu_si256(reinterpret_cast(bp))); \ + __m512i da_hi = _mm512_cvtepi8_epi16( \ + _mm256_loadu_si256(reinterpret_cast(ap + 32))); \ + __m512i db_hi = _mm512_cvtepi8_epi16( \ + _mm256_loadu_si256(reinterpret_cast(bp + 32))); \ + __m512i diff_lo = _mm512_sub_epi16(da_lo, db_lo); \ + __m512i diff_hi = _mm512_sub_epi16(da_hi, db_hi); \ + acc_var = _mm512_dpwssd_epi32(acc_var, diff_lo, diff_lo); \ + acc_var = _mm512_dpwssd_epi32(acc_var, diff_hi, diff_hi); \ + } + EUCL_BLOCK( 0, acc0) + EUCL_BLOCK( 64, acc1) + EUCL_BLOCK(128, acc2) + EUCL_BLOCK(192, acc3) +#undef EUCL_BLOCK + } + __m512i acc = _mm512_add_epi32(_mm512_add_epi32(acc0, acc1), + _mm512_add_epi32(acc2, acc3)); + + // Single 64-byte tail blocks. + for (; i + 64 <= length; i += 64) { + __m512i da_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i))); + __m512i db_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i))); + __m512i da_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i + 32))); + __m512i db_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 32))); + acc = _mm512_dpwssd_epi32(acc, _mm512_sub_epi16(da_lo, db_lo), _mm512_sub_epi16(da_lo, db_lo)); + acc = _mm512_dpwssd_epi32(acc, _mm512_sub_epi16(da_hi, db_hi), _mm512_sub_epi16(da_hi, db_hi)); + } + + // 32-byte tail (one ymm → one 512-bit i16 vector). + if (i + 32 <= length) { + __m512i da = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i))); + __m512i db = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i))); + __m512i diff = _mm512_sub_epi16(da, db); + acc = _mm512_dpwssd_epi32(acc, diff, diff); + i += 32; + } + + int32_t result = hsum_epi32(acc); + + // Scalar tail. + for (; i < length; i++) { + int32_t d = (int32_t)a[i] - (int32_t)b[i]; + result += d * d; + } + return (float)result; +} + +// --------------------------------------------------------------------------- +// cosine_i8 — three parallel VPDPBUSD chains with bias correction +// --------------------------------------------------------------------------- +// +// Computes dot(a,b), ||a||², ||b||² in a single pass using VPDPBUSD. +// Bias trick: a_u = a+128 (unsigned), then subtract 128*sum(b) and 128*sum(a). +// +// dot(a,b) = hsum(VPDPBUSD(acc_dot, a_u, b)) - 128*sum(b) +// ||a||² = hsum(VPDPBUSD(acc_normA, a_u, a)) - 128*sum(a) +// ||b||² = hsum(VPDPBUSD(acc_normB, b_u, b)) - 128*sum(b) +// +// normB reuses the same biasAB accumulator as dot (both need 128*sum(b)). +// 2× unrolled (128 bytes/iteration) with 6 VPDPBUSD + 4 VPDPWSSD per iter. +float cosine_i8(const int8_t * __restrict__ a, size_t aoffset, + const int8_t * __restrict__ b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + + const __m512i bias128 = _mm512_set1_epi16(128); + const __m512i flip = _mm512_set1_epi8(-128); + + __m512i dot0 = _mm512_setzero_si512(), dot1 = _mm512_setzero_si512(); + __m512i normA0 = _mm512_setzero_si512(), normA1 = _mm512_setzero_si512(); + __m512i normB0 = _mm512_setzero_si512(), normB1 = _mm512_setzero_si512(); + __m512i biasAB0 = _mm512_setzero_si512(), biasAB1 = _mm512_setzero_si512(); + __m512i biasA0 = _mm512_setzero_si512(), biasA1 = _mm512_setzero_si512(); + + size_t i = 0; + for (; i + 128 <= length; i += 128) { + __m512i va0 = _mm512_loadu_si512(a + i); + __m512i vb0 = _mm512_loadu_si512(b + i); + __m512i va1 = _mm512_loadu_si512(a + i + 64); + __m512i vb1 = _mm512_loadu_si512(b + i + 64); + __m512i au0 = _mm512_add_epi8(va0, flip); + __m512i bu0 = _mm512_add_epi8(vb0, flip); + __m512i au1 = _mm512_add_epi8(va1, flip); + __m512i bu1 = _mm512_add_epi8(vb1, flip); + + dot0 = _mm512_dpbusd_epi32(dot0, au0, vb0); + dot1 = _mm512_dpbusd_epi32(dot1, au1, vb1); + normA0 = _mm512_dpbusd_epi32(normA0, au0, va0); + normA1 = _mm512_dpbusd_epi32(normA1, au1, va1); + normB0 = _mm512_dpbusd_epi32(normB0, bu0, vb0); + normB1 = _mm512_dpbusd_epi32(normB1, bu1, vb1); + + // Promote to i16 for bias calculations + __m512i va0_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i))); + __m512i va0_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i + 32))); + __m512i vb0_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i))); + __m512i vb0_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 32))); + __m512i va1_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i + 64))); + __m512i va1_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i + 96))); + __m512i vb1_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 64))); + __m512i vb1_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 96))); + + biasAB0 = _mm512_dpwssd_epi32(biasAB0, vb0_lo, bias128); + biasAB0 = _mm512_dpwssd_epi32(biasAB0, vb0_hi, bias128); + biasAB1 = _mm512_dpwssd_epi32(biasAB1, vb1_lo, bias128); + biasAB1 = _mm512_dpwssd_epi32(biasAB1, vb1_hi, bias128); + biasA0 = _mm512_dpwssd_epi32(biasA0, va0_lo, bias128); + biasA0 = _mm512_dpwssd_epi32(biasA0, va0_hi, bias128); + biasA1 = _mm512_dpwssd_epi32(biasA1, va1_lo, bias128); + biasA1 = _mm512_dpwssd_epi32(biasA1, va1_hi, bias128); + } + __m512i dot = _mm512_add_epi32(dot0, dot1); + __m512i normA = _mm512_add_epi32(normA0, normA1); + __m512i normB = _mm512_add_epi32(normB0, normB1); + __m512i biasAB = _mm512_add_epi32(biasAB0, biasAB1); + __m512i biasA = _mm512_add_epi32(biasA0, biasA1); + + // Single-zmm tail. + for (; i + 64 <= length; i += 64) { + __m512i va = _mm512_loadu_si512(a + i); + __m512i vb = _mm512_loadu_si512(b + i); + __m512i au = _mm512_add_epi8(va, flip); + __m512i bu = _mm512_add_epi8(vb, flip); + dot = _mm512_dpbusd_epi32(dot, au, vb); + normA = _mm512_dpbusd_epi32(normA, au, va); + normB = _mm512_dpbusd_epi32(normB, bu, vb); + + // Promote to i16 for bias calculations + __m512i va_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i))); + __m512i va_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(a + i + 32))); + __m512i vb_lo = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i))); + __m512i vb_hi = _mm512_cvtepi8_epi16(_mm256_loadu_si256(reinterpret_cast(b + i + 32))); + + biasAB = _mm512_dpwssd_epi32(biasAB, vb_lo, bias128); + biasAB = _mm512_dpwssd_epi32(biasAB, vb_hi, bias128); + biasA = _mm512_dpwssd_epi32(biasA, va_lo, bias128); + biasA = _mm512_dpwssd_epi32(biasA, va_hi, bias128); + } + + // Apply bias corrections before scalar tail. + int64_t dotResult = (int64_t)hsum_epi32(dot) - (int64_t)hsum_epi32(biasAB); + int64_t normAResult = (int64_t)hsum_epi32(normA) - (int64_t)hsum_epi32(biasA); + int64_t normBResult = (int64_t)hsum_epi32(normB) - (int64_t)hsum_epi32(biasAB); + + // Scalar tail. + for (; i < length; i++) { + int32_t ai = a[i], bi = b[i]; + dotResult += (int64_t)ai * bi; + normAResult += (int64_t)ai * ai; + normBResult += (int64_t)bi * bi; + } + + return (float)(dotResult / sqrt((double)normAResult * (double)normBResult)); +} + } // namespace AVX3_DL diff --git a/jvector-native/src/main/native/src/jvector_simd.cpp b/jvector-native/src/main/native/src/jvector_simd.cpp index 8bf8ebef5..adc6ba3ac 100644 --- a/jvector-native/src/main/native/src/jvector_simd.cpp +++ b/jvector-native/src/main/native/src/jvector_simd.cpp @@ -82,11 +82,14 @@ static const KernelVTable AVX3_vtable = { }; #undef KERNEL_ENTRY -// AVX3_DL (Ice Lake) inherits all slots from AVX3 unchanged for now. -// To override a slot: t.kernel_name = AVX3_DL::kernel_name; -// The implementation must exist in jvector_avx3_dl_kernels.cpp. +// AVX3_DL (Ice Lake) inherits all slots from AVX3, then overrides the three +// int8 similarity kernels with VNNI-accelerated versions from +// jvector_avx3_dl_kernels.cpp. static const KernelVTable AVX3_DL_vtable = []() { KernelVTable t = AVX3_vtable; + t.dot_product_i8 = AVX3_DL::dot_product_i8; + t.euclidean_i8 = AVX3_DL::euclidean_i8; + t.cosine_i8 = AVX3_DL::cosine_i8; return t; }(); diff --git a/jvector-native/src/main/native/src/jvector_simd_kernel_list.h b/jvector-native/src/main/native/src/jvector_simd_kernel_list.h index 63e7baa4d..99bd99245 100644 --- a/jvector-native/src/main/native/src/jvector_simd_kernel_list.h +++ b/jvector-native/src/main/native/src/jvector_simd_kernel_list.h @@ -58,7 +58,11 @@ KERNEL_ENTRY(float, nvq_square_l2_distance_8bit, (const float *vector, const unsigned char *quantized, size_t length, float alpha, float x0, float minValue, float maxValue), (vector, quantized, length, alpha, x0, minValue, maxValue)) \ KERNEL_ENTRY(float, nvq_dot_product_8bit, (const float *vector, const unsigned char *quantized, size_t length, float alpha, float x0, float minValue, float maxValue), (vector, quantized, length, alpha, x0, minValue, maxValue)) \ KERNEL_ENTRY(int64_t, nvq_cosine_8bit_packed, (const float *vector, const unsigned char *quantized, size_t length, float alpha, float x0, float minValue, float maxValue, const float *centroid), (vector, quantized, length, alpha, x0, minValue, maxValue, centroid)) \ - KERNEL_ENTRY(void, nvq_shuffle_query_in_place_8bit, (float *vector, size_t length), (vector, length)) + KERNEL_ENTRY(void, nvq_shuffle_query_in_place_8bit, (float *vector, size_t length), (vector, length)) \ + /* Int8 byte-vector similarity (VNNI-accelerated on AVX3_DL+) */ \ + KERNEL_ENTRY(float, dot_product_i8, (const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length), (a, aoffset, b, boffset, length)) \ + KERNEL_ENTRY(float, euclidean_i8, (const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length), (a, aoffset, b, boffset, length)) \ + KERNEL_ENTRY(float, cosine_i8, (const int8_t *a, size_t aoffset, const int8_t *b, size_t boffset, size_t length), (a, aoffset, b, boffset, length)) /* ── ADD NEW KERNEL_ENTRY LINES ABOVE THIS LINE ── */ // clang-format on diff --git a/jvector-native/src/main/native/src/jvector_simd_kernels.cpp b/jvector-native/src/main/native/src/jvector_simd_kernels.cpp index f4e8c2453..1e13dab55 100644 --- a/jvector-native/src/main/native/src/jvector_simd_kernels.cpp +++ b/jvector-native/src/main/native/src/jvector_simd_kernels.cpp @@ -1640,4 +1640,173 @@ HWY_FLATTEN int64_t nvq_cosine_8bit_packed(const float *HWY_RESTRICT vector, return ((int64_t)bmag_bits << 32) | (int64_t)(uint32_t)sum_bits; } +// ============================================================================= +// Int8 byte-vector similarity kernels +// ============================================================================= +// +// These kernels operate on signed int8 (int8_t) vectors — e.g. the output of +// scalar quantization. The generic path here (compiled for SSE4.2, AVX2, AVX3) +// widens i8→i16 using ReorderWidenMulAccumulate, then accumulates into i32. +// +// On the AVX3_DL (Ice Lake+) tier these implementations are overridden in +// jvector_avx3_dl_kernels.cpp with raw AVX-512 VNNI intrinsics — processing +// 64 bytes per VPDPBUSD clock in a single instruction. +// ============================================================================= + +// Horizontal dot product of two signed int8 vectors. +HWY_FLATTEN float dot_product_i8(const int8_t *HWY_RESTRICT a, size_t aoffset, + const int8_t *HWY_RESTRICT b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + const hn::ScalableTag d8; + const hn::RepartitionToWideX2 d32; // int32, lanes = len(d8)/4 + const hn::Repartition d16; // int16 + const size_t lanes8 = hn::Lanes(d8); + + // Four independent accumulators hide the multi-cycle MADD latency. + auto acc0 = hn::Zero(d32), acc1 = hn::Zero(d32); + auto acc2 = hn::Zero(d32), acc3 = hn::Zero(d32); + auto dummy0 = hn::Zero(d32), dummy1 = hn::Zero(d32); + auto dummy2 = hn::Zero(d32), dummy3 = hn::Zero(d32); + size_t i = 0; + for (; i + 4 * lanes8 <= length; i += 4 * lanes8) { + auto va0 = hn::LoadU(d8, a + i); + auto vb0 = hn::LoadU(d8, b + i); + auto va1 = hn::LoadU(d8, a + i + lanes8); + auto vb1 = hn::LoadU(d8, b + i + lanes8); + auto va2 = hn::LoadU(d8, a + i + 2*lanes8); + auto vb2 = hn::LoadU(d8, b + i + 2*lanes8); + auto va3 = hn::LoadU(d8, a + i + 3*lanes8); + auto vb3 = hn::LoadU(d8, b + i + 3*lanes8); + + // Promote to i16 and accumulate using ReorderWidenMulAccumulate (2 i16s -> 1 i32) + acc0 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va0), hn::PromoteLowerTo(d16, vb0), acc0, dummy0); + acc0 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va0), hn::PromoteUpperTo(d16, vb0), acc0, dummy0); + acc1 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va1), hn::PromoteLowerTo(d16, vb1), acc1, dummy1); + acc1 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va1), hn::PromoteUpperTo(d16, vb1), acc1, dummy1); + acc2 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va2), hn::PromoteLowerTo(d16, vb2), acc2, dummy2); + acc2 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va2), hn::PromoteUpperTo(d16, vb2), acc2, dummy2); + acc3 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va3), hn::PromoteLowerTo(d16, vb3), acc3, dummy3); + acc3 = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va3), hn::PromoteUpperTo(d16, vb3), acc3, dummy3); + } + auto acc = hn::Add(hn::Add(acc0, acc1), hn::Add(acc2, acc3)); + auto dummy = hn::Zero(d32); + + for (; i + lanes8 <= length; i += lanes8) { + auto va = hn::LoadU(d8, a + i); + auto vb = hn::LoadU(d8, b + i); + acc = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va), hn::PromoteLowerTo(d16, vb), acc, dummy); + acc = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va), hn::PromoteUpperTo(d16, vb), acc, dummy); + } + int32_t result = hn::ReduceSum(d32, acc); + for (; i < length; i++) result += (int32_t)a[i] * (int32_t)b[i]; + return (float)result; +} + +// Sum of squared differences of two signed int8 vectors. +// Promote i8→i16, subtract in i16, then ReorderWidenMulAccumulate into i32. +HWY_FLATTEN float euclidean_i8(const int8_t *HWY_RESTRICT a, size_t aoffset, + const int8_t *HWY_RESTRICT b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + const hn::ScalableTag d8; + const hn::RepartitionToWideX2 d32; // int32, lanes = len(d8)/4 + const size_t lanes8 = hn::Lanes(d8); + + auto acc0 = hn::Zero(d32), acc0h = hn::Zero(d32); + auto acc1 = hn::Zero(d32), acc1h = hn::Zero(d32); + auto acc2 = hn::Zero(d32), acc2h = hn::Zero(d32); + auto acc3 = hn::Zero(d32), acc3h = hn::Zero(d32); + size_t i = 0; + for (; i + 4 * lanes8 <= length; i += 4 * lanes8) { +#define DO_EUCL_BLOCK(off, lo_var, hi_var) \ + { \ + const hn::RepartitionToWide _d16; \ + auto _va8 = hn::LoadU(d8, a + i + (off)); \ + auto _vb8 = hn::LoadU(d8, b + i + (off)); \ + auto _diff_lo = hn::Sub(hn::PromoteLowerTo(_d16, _va8), \ + hn::PromoteLowerTo(_d16, _vb8)); \ + auto _diff_hi = hn::Sub(hn::PromoteUpperTo(_d16, _va8), \ + hn::PromoteUpperTo(_d16, _vb8)); \ + lo_var = hn::ReorderWidenMulAccumulate(d32, _diff_lo, _diff_lo, lo_var, hi_var); \ + lo_var = hn::ReorderWidenMulAccumulate(d32, _diff_hi, _diff_hi, lo_var, hi_var); \ + } + DO_EUCL_BLOCK(0, acc0, acc0h) + DO_EUCL_BLOCK(lanes8, acc1, acc1h) + DO_EUCL_BLOCK(2*lanes8, acc2, acc2h) + DO_EUCL_BLOCK(3*lanes8, acc3, acc3h) +#undef DO_EUCL_BLOCK + } + auto acc = hn::Add(hn::Add(acc0, acc1), hn::Add(acc2, acc3)); + auto acch = hn::Add(hn::Add(acc0h, acc1h), hn::Add(acc2h, acc3h)); + acc = hn::Add(acc, acch); + for (; i + lanes8 <= length; i += lanes8) { + const hn::RepartitionToWide d16; + auto va8 = hn::LoadU(d8, a + i); + auto vb8 = hn::LoadU(d8, b + i); + auto diff_lo = hn::Sub(hn::PromoteLowerTo(d16, va8), hn::PromoteLowerTo(d16, vb8)); + auto diff_hi = hn::Sub(hn::PromoteUpperTo(d16, va8), hn::PromoteUpperTo(d16, vb8)); + auto dummy_hi = hn::Zero(d32); + acc = hn::ReorderWidenMulAccumulate(d32, diff_lo, diff_lo, acc, dummy_hi); + acc = hn::Add(acc, dummy_hi); + dummy_hi = hn::Zero(d32); + acc = hn::ReorderWidenMulAccumulate(d32, diff_hi, diff_hi, acc, dummy_hi); + acc = hn::Add(acc, dummy_hi); + } + int32_t result = hn::ReduceSum(d32, acc); + for (; i < length; i++) { + int32_t d = (int32_t)a[i] - (int32_t)b[i]; + result += d * d; + } + return (float)result; +} + +// Cosine similarity of two signed int8 vectors. +// Computes dot(a,b), dot(a,a), dot(b,b) in parallel over a single pass. +HWY_FLATTEN float cosine_i8(const int8_t *HWY_RESTRICT a, size_t aoffset, + const int8_t *HWY_RESTRICT b, size_t boffset, + size_t length) +{ + a += aoffset; + b += boffset; + const hn::ScalableTag d8; + const hn::RepartitionToWideX2 d32; + const hn::Repartition d16; + const size_t lanes8 = hn::Lanes(d8); + + auto dot = hn::Zero(d32); + auto normA = hn::Zero(d32); + auto normB = hn::Zero(d32); + auto dummy_acc = hn::Zero(d32); + + size_t i = 0; + for (; i + lanes8 <= length; i += lanes8) { + auto va = hn::LoadU(d8, a + i); + auto vb = hn::LoadU(d8, b + i); + + dot = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va), hn::PromoteLowerTo(d16, vb), dot, dummy_acc); + dot = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va), hn::PromoteUpperTo(d16, vb), dot, dummy_acc); + normA = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, va), hn::PromoteLowerTo(d16, va), normA, dummy_acc); + normA = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, va), hn::PromoteUpperTo(d16, va), normA, dummy_acc); + normB = hn::ReorderWidenMulAccumulate(d32, hn::PromoteLowerTo(d16, vb), hn::PromoteLowerTo(d16, vb), normB, dummy_acc); + normB = hn::ReorderWidenMulAccumulate(d32, hn::PromoteUpperTo(d16, vb), hn::PromoteUpperTo(d16, vb), normB, dummy_acc); + } + + int64_t dotResult = (int64_t)hn::ReduceSum(d32, dot); + int64_t normAResult = (int64_t)hn::ReduceSum(d32, normA); + int64_t normBResult = (int64_t)hn::ReduceSum(d32, normB); + + for (; i < length; i++) { + int32_t ai = a[i], bi = b[i]; + dotResult += ai * bi; + normAResult += ai * ai; + normBResult += bi * bi; + } + return (float)(dotResult / sqrt((double)normAResult * (double)normBResult)); +} + } // namespace JV_ISA diff --git a/jvector-native/src/main/native/tests/test_helpers.cpp b/jvector-native/src/main/native/tests/test_helpers.cpp index 957e42947..45b35cc6e 100644 --- a/jvector-native/src/main/native/tests/test_helpers.cpp +++ b/jvector-native/src/main/native/tests/test_helpers.cpp @@ -73,6 +73,10 @@ const std::vector kKernelTestParams = { {100, "large_mixed_tail"}, {128, "large_power_of_2"}, {255, "large_odd_tail_15"}, + // ---- i8 VNNI 256-byte unroll boundaries (dot_product_i8/euclidean_i8) - + {256, "i8_vnni_4x_exact"}, + {263, "i8_vnni_4x_tail_7"}, + {135, "i8_vnni_2zmm_tail_7"}, }; std::vector make_vec(size_t n, float seed) @@ -86,3 +90,21 @@ std::vector make_vec(size_t n, float seed) return v; } +// Produces n int8_t values with a mix of signs and magnitudes. +// The pattern ensures no element is zero (important for cosine tests). +std::vector make_vec_i8(size_t n, int8_t seed) +{ + std::vector v(n); + for (size_t i = 0; i < n; ++i) { + // Scale seed by a small per-element factor to get variety, + // then clamp to [-100, 100] to keep products well within int16 range. + int val = static_cast(seed) + static_cast(i % 13) - 6; + if (i % 3 == 0) val = -val; // mix of signs + if (val == 0) val = 1; // never zero + if (val > 100) val = 100; + if (val < -100) val = -100; + v[i] = static_cast(val); + } + return v; +} + diff --git a/jvector-native/src/main/native/tests/test_helpers.h b/jvector-native/src/main/native/tests/test_helpers.h index a48ea47cf..e1e9cd79c 100644 --- a/jvector-native/src/main/native/tests/test_helpers.h +++ b/jvector-native/src/main/native/tests/test_helpers.h @@ -35,7 +35,11 @@ // so that no element is exactly zero (important for cosine tests). // --------------------------------------------------------------------------- -std::vector make_vec(size_t n, float seed); +std::vector make_vec(size_t n, float seed); + +// make_vec_i8(n, seed) produces n int8_t values with a mix of signs +// suitable for testing the i8 similarity kernels. +std::vector make_vec_i8(size_t n, int8_t seed); // --------------------------------------------------------------------------- // Shared test parameter — vector length + human-readable path description. diff --git a/jvector-native/src/main/native/tests/test_similarity_i8.cpp b/jvector-native/src/main/native/tests/test_similarity_i8.cpp new file mode 100644 index 000000000..51443a4bd --- /dev/null +++ b/jvector-native/src/main/native/tests/test_similarity_i8.cpp @@ -0,0 +1,229 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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. + */ + +// Tests for int8 vector similarity kernels: dot_product_i8, euclidean_i8, cosine_i8. +// +// The kernels operate on signed int8_t vectors and return a float result: +// dot_product_i8 — (float) sum(a[i] * b[i]) +// euclidean_i8 — (float) sum((a[i] - b[i])^2) (squared L2 distance) +// cosine_i8 — (float) dot(a,b) / sqrt(||a||^2 * ||b||^2) +// +// On AVX3_DL (Ice Lake+) these are overridden with VNNI (VPDPBUSD/VPDPWSSD) +// implementations; on all other tiers the generic Highway path is used. +// +// All tests are parametrised over kKernelTestParams (defined in test_helpers.cpp), +// which covers every ISA-tier loop-boundary for both f32 and i8 kernels, including +// the VNNI-specific 64/128/256-byte unroll boundaries added for the i8 suite. + +#include "test_helpers.h" + +// --------------------------------------------------------------------------- +// Reference scalar implementations +// --------------------------------------------------------------------------- + +static float ref_dot_i8(const std::vector& a, const std::vector& b) +{ + int64_t s = 0; + for (size_t i = 0; i < a.size(); ++i) + s += static_cast(a[i]) * static_cast(b[i]); + return static_cast(s); +} + +static float ref_euclidean_i8(const std::vector& a, const std::vector& b) +{ + int64_t s = 0; + for (size_t i = 0; i < a.size(); ++i) { + int32_t d = static_cast(a[i]) - static_cast(b[i]); + s += d * d; + } + return static_cast(s); +} + +static float ref_cosine_i8(const std::vector& a, const std::vector& b) +{ + int64_t dot = 0, normA = 0, normB = 0; + for (size_t i = 0; i < a.size(); ++i) { + int32_t ai = a[i], bi = b[i]; + dot += static_cast(ai) * bi; + normA += static_cast(ai) * ai; + normB += static_cast(bi) * bi; + } + return static_cast(dot / std::sqrt(static_cast(normA) + * static_cast(normB))); +} + +// --------------------------------------------------------------------------- +// Parametrised test fixture +// --------------------------------------------------------------------------- + +class SimilarityI8Test : public ::testing::TestWithParam {}; + +// --------------------------------------------------------------------------- +// dot_product_i8 — SIMD result must match the scalar reference +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, DotProduct) +{ + const size_t n = GetParam().length; + auto a = make_vec_i8(n, 7); + auto b = make_vec_i8(n, 11); + + const float want = ref_dot_i8(a, b); + const float got = dot_product_i8(a.data(), 0, b.data(), 0, n); + + // Integer accumulation with a single int64→float cast — result is exact. + EXPECT_EQ(got, want); +} + +// --------------------------------------------------------------------------- +// dot_product_i8 with non-zero offsets — exercises the aoffset/boffset path +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, DotProductWithOffset) +{ + const size_t n = GetParam().length; + const size_t prefix = 5; // arbitrary prefix that must be ignored + + std::vector a_pad(prefix + n, 0); + std::vector b_pad(prefix + n, 0); + auto a = make_vec_i8(n, 7); + auto b = make_vec_i8(n, 11); + std::copy(a.begin(), a.end(), a_pad.begin() + prefix); + std::copy(b.begin(), b.end(), b_pad.begin() + prefix); + + const float want = ref_dot_i8(a, b); + const float got = dot_product_i8(a_pad.data(), prefix, b_pad.data(), prefix, n); + + EXPECT_EQ(got, want); +} + +// --------------------------------------------------------------------------- +// dot_product_i8 — zero vector gives exactly 0.0 +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, DotProductZeroVector) +{ + const size_t n = GetParam().length; + auto a = make_vec_i8(n, 7); + std::vector z(n, 0); + + EXPECT_EQ(dot_product_i8(a.data(), 0, z.data(), 0, n), 0.0f); +} + +// --------------------------------------------------------------------------- +// euclidean_i8 — SIMD result must match the scalar reference +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, Euclidean) +{ + const size_t n = GetParam().length; + auto a = make_vec_i8(n, 7); + auto b = make_vec_i8(n, 11); + + const float want = ref_euclidean_i8(a, b); + const float got = euclidean_i8(a.data(), 0, b.data(), 0, n); + + // Integer accumulation with a single int64→float cast — result is exact. + EXPECT_EQ(got, want); +} + +// --------------------------------------------------------------------------- +// euclidean_i8 — identical vectors must give exactly 0 +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, EuclideanSameVector) +{ + const size_t n = GetParam().length; + auto a = make_vec_i8(n, 9); + + const float got = euclidean_i8(a.data(), 0, a.data(), 0, n); + + EXPECT_EQ(got, 0.0f); +} + +// --------------------------------------------------------------------------- +// cosine_i8 — SIMD result must match the scalar reference +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, Cosine) +{ + const size_t n = GetParam().length; + auto a = make_vec_i8(n, 7); + auto b = make_vec_i8(n, 11); + + const float want = ref_cosine_i8(a, b); + const float got = cosine_i8(a.data(), 0, b.data(), 0, n); + + EXPECT_NEAR(got, want, 1e-5f); +} + +// --------------------------------------------------------------------------- +// cosine_i8 — parallel vectors (b = k*a, k > 0) should give similarity ≈ 1.0 +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, CosineParallelVectors) +{ + const size_t n = GetParam().length; + // Use small magnitudes so that 2*val stays within int8 range. + auto a = make_vec_i8(n, 3); + std::vector b(n); + for (size_t i = 0; i < n; ++i) + b[i] = static_cast(std::max(-127, std::min(127, 2 * static_cast(a[i])))); + + const float got = cosine_i8(a.data(), 0, b.data(), 0, n); + + EXPECT_NEAR(got, 1.0f, 1e-5f); +} + +// --------------------------------------------------------------------------- +// cosine_i8 — orthogonal vectors should give similarity ≈ 0.0 +// +// Same analytic construction as the f32 test: for even n, +// a = [+1, +1, +1, ...] +// b = [+1, -1, +1, -1, ...] → dot(a,b) = 0. +// Odd n: the odd last element is zeroed out on b (unchanged on a) so the +// dot product remains zero without affecting the norms materially. +// --------------------------------------------------------------------------- + +TEST_P(SimilarityI8Test, CosineOrthogonalVectors) +{ + const size_t n = GetParam().length; + if (n < 2) GTEST_SKIP() << "need at least 2 elements for orthogonality"; + + const size_t even_n = n - (n % 2); + + std::vector a(n, 0), b(n, 0); + for (size_t i = 0; i < even_n; ++i) { + a[i] = 1; + b[i] = (i % 2 == 0) ? 1 : -1; + } + + const float got = cosine_i8(a.data(), 0, b.data(), 0, n); + + EXPECT_NEAR(got, 0.0f, 1e-5f); +} + +// --------------------------------------------------------------------------- +// Instantiation — named using the description field +// --------------------------------------------------------------------------- + +INSTANTIATE_TEST_SUITE_P( + AllSizes, + SimilarityI8Test, + ::testing::ValuesIn(kKernelTestParams), + [](const ::testing::TestParamInfo& info) { + return info.param.description; + }); diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java index 81a99aafc..29f0ec18f 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java @@ -19,6 +19,7 @@ import com.carrotsearch.randomizedtesting.RandomizedTest; import io.github.jbellis.jvector.TestUtil; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.FloatArray; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; @@ -60,6 +61,39 @@ public void testSimilarityMetricsFloat() { Assert.assertEquals(a.getVectorUtilSupport().squareDistance(v1a, v2a), b.getVectorUtilSupport().squareDistance(v1b, v2b), 0.0001f); } + @Test + public void testSimilarityMetricsByte() { + Assume.assumeTrue(hasSimd); + + VectorizationProvider a = new DefaultVectorizationProvider(); + VectorizationProvider b = VectorizationProvider.getInstance(); + + // Use a prime-length vector that is not a multiple of 8 or 16 + int dim = 107; + byte[] rawA = new byte[dim]; + byte[] rawB = new byte[dim]; + getRandom().nextBytes(rawA); + getRandom().nextBytes(rawB); + + ByteSequence bsA_scalar = a.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence bsB_scalar = a.getVectorTypeSupport().createByteSequence(rawB); + ByteSequence bsA_simd = b.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence bsB_simd = b.getVectorTypeSupport().createByteSequence(rawB); + + Assert.assertEquals( + a.getVectorUtilSupport().dotProduct(bsA_scalar, bsB_scalar), + b.getVectorUtilSupport().dotProduct(bsA_simd, bsB_simd), + 0.0001f); + Assert.assertEquals( + a.getVectorUtilSupport().squareDistance(bsA_scalar, bsB_scalar), + b.getVectorUtilSupport().squareDistance(bsA_simd, bsB_simd), + 0.0001f); + Assert.assertEquals( + a.getVectorUtilSupport().cosine(bsA_scalar, bsB_scalar), + b.getVectorUtilSupport().cosine(bsA_simd, bsB_simd), + 0.0001f); + } + @Test public void testAssembleAndSum() { Assume.assumeTrue(hasSimd); diff --git a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java index f95d4bf6b..df49f8858 100644 --- a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java +++ b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java @@ -976,43 +976,266 @@ float assembleAndSumPQ_512( return res; } + // ----------------------------------------------------------------------- + // ByteSequence similarity metrics – Panama SIMD implementations + // + // Strategy: widen signed bytes to int32 via B2I (no AND-mask needed for + // signed arithmetic), accumulate products in IntVector lanes, then reduce. + // The byte-vector species is 1/4 the width of the int species: + // 512-bit int (16 lanes) <- SPECIES_128 bytes + // 256-bit int (8 lanes) <- SPECIES_64 bytes + // 128-bit preferred <- scalar (ByteVector.SPECIES_32 does not exist; + // 128-bit SIMD shows no benefit for this workload) + // ----------------------------------------------------------------------- + /** - * Vectorized calculation of Hamming distance for two arrays of long integers. - * Both arrays should have the same length. - * - * @param a The first array - * @param b The second array - * @return The Hamming distance + * Vectorized dot product of two signed int8 byte vectors. */ @Override public float dotProduct(ByteSequence a, ByteSequence b) { - float sum = 0; - for (int i = 0; i < a.length(); i++) { - sum += (int) a.get(i) * (int) b.get(i); + return switch (PREFERRED_BIT_SIZE) { + case 512 -> dotProductBytes512(a, b); + case 256 -> dotProductBytes256(a, b); + default -> dotProductBytes128(a, b); + }; + } + + float dotProductBytes512(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_128.length(); // 16 + final int limit = ByteVector.SPECIES_128.loopBound(length); + IntVector acc = IntVector.zero(IntVector.SPECIES_512); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + acc = acc.add(va.mul(vb)); } - return sum; + + int result = acc.reduceLanes(VectorOperators.ADD); + for (int i = limit; i < length; i++) { + result += a.get(aOff + i) * b.get(bOff + i); + } + return result; + } + + float dotProductBytes256(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_64.length(); // 8 + final int limit = ByteVector.SPECIES_64.loopBound(length); + IntVector acc = IntVector.zero(IntVector.SPECIES_256); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + acc = acc.add(va.mul(vb)); + } + + int result = acc.reduceLanes(VectorOperators.ADD); + for (int i = limit; i < length; i++) { + result += a.get(aOff + i) * b.get(bOff + i); + } + return result; + } + + float dotProductBytes128(ByteSequence a, ByteSequence b) { + // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + int result = 0; + for (int i = 0; i < length; i++) { + result += a.get(aOff + i) * b.get(bOff + i); + } + return result; } + /** + * Vectorized sum of squared differences between two signed int8 byte vectors. + */ @Override public float squareDistance(ByteSequence a, ByteSequence b) { - float sum = 0; - for (int i = 0; i < a.length(); i++) { - float diff = a.get(i) - b.get(i); - sum += diff * diff; + return switch (PREFERRED_BIT_SIZE) { + case 512 -> squareDistanceBytes512(a, b); + case 256 -> squareDistanceBytes256(a, b); + default -> squareDistanceBytes128(a, b); + }; + } + + float squareDistanceBytes512(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_128.length(); + final int limit = ByteVector.SPECIES_128.loopBound(length); + IntVector acc = IntVector.zero(IntVector.SPECIES_512); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector diff = va.sub(vb); + acc = acc.add(diff.mul(diff)); } - return sum; + + int result = acc.reduceLanes(VectorOperators.ADD); + for (int i = limit; i < length; i++) { + int diff = a.get(aOff + i) - b.get(bOff + i); + result += diff * diff; + } + return result; + } + + float squareDistanceBytes256(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_64.length(); + final int limit = ByteVector.SPECIES_64.loopBound(length); + IntVector acc = IntVector.zero(IntVector.SPECIES_256); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector diff = va.sub(vb); + acc = acc.add(diff.mul(diff)); + } + + int result = acc.reduceLanes(VectorOperators.ADD); + for (int i = limit; i < length; i++) { + int diff = a.get(aOff + i) - b.get(bOff + i); + result += diff * diff; + } + return result; + } + + float squareDistanceBytes128(ByteSequence a, ByteSequence b) { + // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + int result = 0; + for (int i = 0; i < length; i++) { + int diff = a.get(aOff + i) - b.get(bOff + i); + result += diff * diff; + } + return result; } + /** + * Vectorized cosine similarity between two signed int8 byte vectors. + */ @Override public float cosine(ByteSequence a, ByteSequence b) { - float dot = 0, normA = 0, normB = 0; - for (int i = 0; i < a.length(); i++) { - float ai = a.get(i), bi = b.get(i); - dot += ai * bi; - normA += ai * ai; - normB += bi * bi; - } - return (float) (dot / Math.sqrt(normA * normB)); + return switch (PREFERRED_BIT_SIZE) { + case 512 -> cosineBytes512(a, b); + case 256 -> cosineBytes256(a, b); + default -> cosineBytes128(a, b); + }; + } + + float cosineBytes512(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_128.length(); + final int limit = ByteVector.SPECIES_128.loopBound(length); + IntVector dot = IntVector.zero(IntVector.SPECIES_512); + IntVector normA = IntVector.zero(IntVector.SPECIES_512); + IntVector normB = IntVector.zero(IntVector.SPECIES_512); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + dot = dot.add(va.mul(vb)); + normA = normA.add(va.mul(va)); + normB = normB.add(vb.mul(vb)); + } + + long dotResult = dot.reduceLanes(VectorOperators.ADD); + long normAResult = normA.reduceLanes(VectorOperators.ADD); + long normBResult = normB.reduceLanes(VectorOperators.ADD); + + for (int i = limit; i < length; i++) { + int ai = a.get(aOff + i), bi = b.get(bOff + i); + dotResult += ai * bi; + normAResult += ai * ai; + normBResult += bi * bi; + } + return (float) (dotResult / Math.sqrt((double) normAResult * normBResult)); + } + + float cosineBytes256(ByteSequence a, ByteSequence b) { + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + final int step = ByteVector.SPECIES_64.length(); + final int limit = ByteVector.SPECIES_64.loopBound(length); + IntVector dot = IntVector.zero(IntVector.SPECIES_256); + IntVector normA = IntVector.zero(IntVector.SPECIES_256); + IntVector normB = IntVector.zero(IntVector.SPECIES_256); + + for (int i = 0; i < limit; i += step) { + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + dot = dot.add(va.mul(vb)); + normA = normA.add(va.mul(va)); + normB = normB.add(vb.mul(vb)); + } + + long dotResult = dot.reduceLanes(VectorOperators.ADD); + long normAResult = normA.reduceLanes(VectorOperators.ADD); + long normBResult = normB.reduceLanes(VectorOperators.ADD); + + for (int i = limit; i < length; i++) { + int ai = a.get(aOff + i), bi = b.get(bOff + i); + dotResult += ai * bi; + normAResult += ai * ai; + normBResult += bi * bi; + } + return (float) (dotResult / Math.sqrt((double) normAResult * normBResult)); + } + + float cosineBytes128(ByteSequence a, ByteSequence b) { + // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width + final int length = a.length(); + final int aOff = a.offset(); + final int bOff = b.offset(); + long dotResult = 0, normAResult = 0, normBResult = 0; + for (int i = 0; i < length; i++) { + int ai = a.get(aOff + i), bi = b.get(bOff + i); + dotResult += ai * bi; + normAResult += ai * ai; + normBResult += bi * bi; + } + return (float) (dotResult / Math.sqrt((double) normAResult * normBResult)); } @Override From 2685d3236a2f7c338aa68ac5b85e06d07d3f92e2 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Mon, 24 Aug 2026 06:59:22 +0000 Subject: [PATCH 07/13] Add InlineByteVectors feature for native int8 on-disk storage MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Introduces InlineByteVectors, a new Feature that stores signed int8 vectors inline in an OnDiskGraphIndex at 1 byte per component — 4x more compact than the float32 InlineVectors representation. Changes: - InlineByteVectors.java: new Feature implementation backed by VectorTypeSupport.writeByteSequence / readByteSequence - FeatureId: add INLINE_BYTE_VECTORS as ordinal 5 (backward-compatible) - AbstractGraphIndexWriter.Builder: accept INLINE_BYTE_VECTORS as the canonical source for the vector dimension in the file header - OnDiskGraphIndex.View: add getByteVector(int) to read a stored int8 vector from disk, and byteVectorRerankerFor(ByteSequence, ByteVectorSimilarityFunction) to wire byte-by-byte disk scoring directly into the search path --- .../graph/disk/AbstractGraphIndexWriter.java | 4 + .../jvector/graph/disk/OnDiskGraphIndex.java | 38 ++++++++ .../jvector/graph/disk/feature/FeatureId.java | 3 +- .../graph/disk/feature/InlineByteVectors.java | 90 +++++++++++++++++++ .../graph/disk/TestOnDiskGraphIndex.java | 2 +- 5 files changed, 135 insertions(+), 2 deletions(-) create mode 100644 jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/InlineByteVectors.java diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexWriter.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexWriter.java index 305b0850f..a68ab9bc0 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexWriter.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexWriter.java @@ -20,6 +20,8 @@ import io.github.jbellis.jvector.graph.ImmutableGraphIndex; import io.github.jbellis.jvector.graph.disk.feature.Feature; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; +import io.github.jbellis.jvector.graph.disk.feature.FusedFeature; +import io.github.jbellis.jvector.graph.disk.feature.InlineByteVectors; import io.github.jbellis.jvector.graph.disk.feature.InlineVectors; import io.github.jbellis.jvector.graph.disk.feature.NVQ; import io.github.jbellis.jvector.graph.disk.feature.SeparatedFeature; @@ -259,6 +261,8 @@ public K build() throws IOException { int dimension; if (features.containsKey(FeatureId.INLINE_VECTORS)) { dimension = ((InlineVectors) features.get(FeatureId.INLINE_VECTORS)).dimension(); + } else if (features.containsKey(FeatureId.INLINE_BYTE_VECTORS)) { + dimension = ((InlineByteVectors) features.get(FeatureId.INLINE_BYTE_VECTORS)).dimension(); } else if (features.containsKey(FeatureId.NVQ_VECTORS)) { dimension = ((NVQ) features.get(FeatureId.NVQ_VECTORS)).dimension(); } else if (features.containsKey(FeatureId.SEPARATED_VECTORS)) { diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/OnDiskGraphIndex.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/OnDiskGraphIndex.java index f43357c27..1877cce04 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/OnDiskGraphIndex.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/OnDiskGraphIndex.java @@ -27,6 +27,7 @@ import io.github.jbellis.jvector.graph.disk.feature.FeatureSource; import io.github.jbellis.jvector.graph.disk.feature.FusedPQ; import io.github.jbellis.jvector.graph.disk.feature.FusedFeature; +import io.github.jbellis.jvector.graph.disk.feature.InlineByteVectors; import io.github.jbellis.jvector.graph.disk.feature.InlineVectors; import io.github.jbellis.jvector.graph.disk.feature.NVQ; import io.github.jbellis.jvector.graph.disk.feature.SeparatedFeature; @@ -36,8 +37,10 @@ import java.util.ArrayList; import io.github.jbellis.jvector.util.Bits; import io.github.jbellis.jvector.util.RamUsageEstimator; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; @@ -578,6 +581,41 @@ public void getVectorInto(int node, VectorFloat vector, int offset) { } } + /** + * Returns the signed int8 vector stored for {@code node} via the + * {@link FeatureId#INLINE_BYTE_VECTORS} feature. + * + * @throws UnsupportedOperationException if the graph was not written with + * {@link InlineByteVectors} + */ + public ByteSequence getByteVector(int node) { + if (!features.containsKey(FeatureId.INLINE_BYTE_VECTORS)) { + throw new UnsupportedOperationException("No inline byte vectors in this graph"); + } + try { + long diskOffset = offsetFor(node, FeatureId.INLINE_BYTE_VECTORS); + reader.seek(diskOffset); + return vectorTypeSupport.readByteSequence(reader, dimension); + } catch (IOException e) { + throw new UncheckedIOException(e); + } + } + + /** + * Returns a {@link ScoreFunction.ExactScoreFunction} that scores candidates by reading + * their int8 vectors from disk and comparing them byte×byte against {@code queryBytes}. + * + * @throws UnsupportedOperationException if the graph was not written with + * {@link InlineByteVectors} + */ + public ScoreFunction.ExactScoreFunction byteVectorRerankerFor(ByteSequence queryBytes, + ByteVectorSimilarityFunction bvsf) { + if (!features.containsKey(FeatureId.INLINE_BYTE_VECTORS)) { + throw new UnsupportedOperationException("No inline byte vectors in this graph"); + } + return node -> bvsf.compare(queryBytes, getByteVector(node)); + } + public NodesIterator getNeighborsIterator(int level, int node) { try { int[] stored; diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/FeatureId.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/FeatureId.java index 131c4c8ee..37b62e43a 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/FeatureId.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/FeatureId.java @@ -33,7 +33,8 @@ public enum FeatureId { FUSED_PQ(FusedPQ::load), NVQ_VECTORS(NVQ::load), SEPARATED_VECTORS(SeparatedVectors::load), - SEPARATED_NVQ(SeparatedNVQ::load); + SEPARATED_NVQ(SeparatedNVQ::load), + INLINE_BYTE_VECTORS(InlineByteVectors::load); private final BiFunction loader; diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/InlineByteVectors.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/InlineByteVectors.java new file mode 100644 index 000000000..a51e8398e --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/feature/InlineByteVectors.java @@ -0,0 +1,90 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.graph.disk.feature; + +import io.github.jbellis.jvector.disk.IndexWriter; +import io.github.jbellis.jvector.disk.RandomAccessReader; +import io.github.jbellis.jvector.graph.disk.CommonHeader; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; + +import java.io.IOException; + +/** + * Stores signed int8 (byte) vectors inline in an {@link io.github.jbellis.jvector.graph.disk.OnDiskGraphIndex}. + *

+ * Each vector occupies exactly {@code dimension} bytes on disk — 4× smaller than + * the float32 {@link InlineVectors} representation. The on-disk layout is otherwise + * identical: one contiguous byte block per node record, written and read via + * {@link VectorTypeSupport#writeByteSequence} and {@link VectorTypeSupport#readByteSequence}. + *

+ * Use {@link io.github.jbellis.jvector.graph.disk.OnDiskGraphIndex.View#getByteVector(int)} + * to retrieve a stored vector at search time. + */ +public class InlineByteVectors extends AbstractFeature { + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + private final int dimension; + + public InlineByteVectors(int dimension) { + this.dimension = dimension; + } + + @Override + public FeatureId id() { + return FeatureId.INLINE_BYTE_VECTORS; + } + + /** No extra header bytes — dimension is already in the {@link CommonHeader}. */ + @Override + public int headerSize() { + return 0; + } + + /** One byte per component. */ + @Override + public int featureSize() { + return dimension; + } + + public int dimension() { + return dimension; + } + + static InlineByteVectors load(CommonHeader header, RandomAccessReader reader) { + return new InlineByteVectors(header.dimension); + } + + @Override + public void writeHeader(IndexWriter out) { + // common header carries dimension; nothing extra needed + } + + @Override + public void writeInline(IndexWriter out, Feature.State state) throws IOException { + vts.writeByteSequence(out, ((State) state).vector); + } + + public static class State implements Feature.State { + public final ByteSequence vector; + + public State(ByteSequence vector) { + this.vector = vector; + } + } +} diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java index c76796bb4..c84d5b103 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java @@ -340,7 +340,7 @@ public void testV0Read() throws IOException { var onDiskGraph = OnDiskGraphIndex.load(readerSupplier); var onDiskView = onDiskGraph.getView()) { - assertEquals(88, onDiskGraph.ramBytesUsed()); // Current size of graph that hasn't been read from + assertEquals(96, onDiskGraph.ramBytesUsed()); // Current size of graph that hasn't been read from assertEquals(32, onDiskGraph.getDegree(0)); assertEquals(2, onDiskGraph.version); assertEquals(100_000, onDiskGraph.size(0)); From 62dcebe33964d06e004cd57bb8f0a4d5465e1f9a Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Mon, 24 Aug 2026 07:03:23 +0000 Subject: [PATCH 08/13] Add INT8 end-to-end tutorial (Int8Example) Demonstrates the full int8 pipeline using the siftsmall dataset: - Read siftsmall_base.fvecs and convert float32 vectors to signed int8 - Build a graph index with byte-by-byte scoring (no float32 round-trip) - Save the graph to disk using InlineByteVectors (1 byte/component) - Load the index from disk - Search with random int8 query vectors, scoring directly from disk byte vectors Also registers the tutorial under the 'int8' key in TutorialRunner. --- .../jvector/example/tutorial/Int8Example.java | 204 ++++++++++++++++++ .../example/tutorial/TutorialRunner.java | 3 + 2 files changed, 207 insertions(+) create mode 100644 jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/Int8Example.java diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/Int8Example.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/Int8Example.java new file mode 100644 index 000000000..8678b8b03 --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/Int8Example.java @@ -0,0 +1,204 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.example.tutorial; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Random; + +import io.github.jbellis.jvector.disk.ReaderSupplier; +import io.github.jbellis.jvector.disk.ReaderSupplierFactory; +import io.github.jbellis.jvector.example.util.SiftLoader; +import io.github.jbellis.jvector.graph.GraphIndexBuilder; +import io.github.jbellis.jvector.graph.GraphSearcher; +import io.github.jbellis.jvector.graph.ImmutableGraphIndex; +import io.github.jbellis.jvector.graph.ListRandomAccessByteVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessByteVectorValues; +import io.github.jbellis.jvector.graph.SearchResult; +import io.github.jbellis.jvector.graph.disk.GraphIndexWriter; +import io.github.jbellis.jvector.graph.disk.GraphIndexWriterTypes; +import io.github.jbellis.jvector.graph.disk.OnDiskGraphIndex; +import io.github.jbellis.jvector.graph.disk.feature.FeatureId; +import io.github.jbellis.jvector.graph.disk.feature.InlineByteVectors; +import io.github.jbellis.jvector.graph.similarity.DefaultSearchScoreProvider; +import io.github.jbellis.jvector.graph.similarity.SearchScoreProvider; +import io.github.jbellis.jvector.util.Bits; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; + +/** + * Demonstrates end-to-end INT8 (signed byte) vector support in JVector: + * + *

How to run: + *

+ *   # 1. Download the siftsmall dataset from http://corpus-texmex.irisa.fr/
+ *   #    and unzip it so that siftsmall/siftsmall_base.fvecs is present
+ *   #    in the directory you run the command from.
+ *
+ *   # 2. Build and run — jdk22 is the default profile (active when no -P flag is given):
+ *   ./mvnw compile -pl jvector-examples -am
+ *   ./mvnw exec:exec@tutorial -pl jvector-examples -Dtutorial=int8
+ *
+ *   # Use -Pjdk20 or -Pjdk11 if you are on an older JDK:
+ *   ./mvnw exec:exec@tutorial -pl jvector-examples -Pjdk20 -Dtutorial=int8
+ * 
+ * + *
    + *
  1. Read the siftsmall dataset from .fvecs files (float32).
  2. + *
  3. Convert each float32 vector to a signed int8 {@link ByteSequence}.
  4. + *
  5. Build a graph index using only byte×byte distance calculations.
  6. + *
  7. Save the graph to disk with {@link InlineByteVectors} — 1 byte per component on disk.
  8. + *
  9. Load the index back from disk.
  10. + *
  11. Generate random int8 query vectors and search, scoring directly from disk byte vectors.
  12. + *
+ * + *

The siftsmall dataset must be present at {@code siftsmall/siftsmall_base.fvecs} + * relative to the working directory. Download it from + * http://corpus-texmex.irisa.fr/. + */ +public class Int8Example { + + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + public static void main(String[] args) throws IOException { + String siftDir = args.length > 0 ? args[0] : "siftsmall"; + + // ── Step 1: Read siftsmall base vectors (.fvecs) ───────────────────────────── + System.out.println("Loading siftsmall base vectors..."); + List> floatVectors = SiftLoader.readFvecs(siftDir + "/siftsmall_base.fvecs"); + int dimension = floatVectors.get(0).length(); + System.out.printf("Loaded %d vectors of dimension %d%n", floatVectors.size(), dimension); + + // ── Step 2: Convert float32 vectors to signed int8 (ByteSequence) ──────────── + // SIFT base vectors store gradient histogram bins as unsigned bytes in [0, 255]. + // We subtract 128 to shift them into the signed range [-128, 127] that + // ByteVectorSimilarityFunction and the underlying SIMD routines expect. + // For other float32 datasets you would typically scale by a dataset-specific + // factor and then clamp before casting. + System.out.println("Converting float32 vectors to int8..."); + List> byteVectors = new ArrayList<>(floatVectors.size()); + for (VectorFloat fv : floatVectors) { + ByteSequence bv = vts.createByteSequence(dimension); + for (int i = 0; i < dimension; i++) { + // SIFT component in [0,255] → shift to signed [-128, 127] + bv.set(i, (byte) ((int) fv.get(i) - 128)); + } + byteVectors.add(bv); + } + + // Wrap the list in a RandomAccessByteVectorValues (RABVV) — + // the byte-vector analogue of RandomAccessVectorValues. + RandomAccessByteVectorValues rabvv = new ListRandomAccessByteVectorValues(byteVectors, dimension); + + // ── Step 3: Build the graph index using int8 vectors ───────────────────────── + // Use the fluent builder API. addHierarchy and refineFinalGraph are controlled + // via GraphIndexBuilderConfig (JMX/system property); no need to pass them here. + // The BuildScoreProvider wires up byte×byte scoring — no float32 round-trip + // occurs during construction. + int M = 32; + int efConstruction = 100; + float neighborOverflow = 1.2f; + float alpha = 1.2f; + + System.out.println("Building graph index from int8 vectors..."); + ImmutableGraphIndex heapGraph; + try (GraphIndexBuilder builder = GraphIndexBuilder.builder(rabvv, ByteVectorSimilarityFunction.EUCLIDEAN, M) + .withBeamWidth(efConstruction) + .withNeighborOverflow(neighborOverflow) + .withAlpha(alpha) + .build()) + { + heapGraph = builder.build(rabvv); + } + System.out.printf("Graph built with %d nodes%n", heapGraph.size(0)); + + // ── Step 4: Save the graph to disk with native int8 storage ────────────────── + // InlineByteVectors stores each vector as `dimension` raw bytes on disk — + // 4× more compact than the float32 InlineVectors alternative. + Path graphPath = Files.createTempFile("jvector-int8-example", null); + System.out.printf("Writing graph to disk (%d bytes/vector): %s%n", dimension, graphPath); + try (GraphIndexWriter writer = GraphIndexWriter + .getBuilderFor(GraphIndexWriterTypes.RANDOM_ACCESS_PARALLEL, heapGraph, graphPath) + .with(new InlineByteVectors(dimension)) + .build()) + { + writer.write(Map.of( + FeatureId.INLINE_BYTE_VECTORS, + nodeId -> new InlineByteVectors.State(rabvv.getVector(nodeId)) + )); + } + + // ── Step 5: Load the index from disk ───────────────────────────────────────── + System.out.println("Loading graph from disk..."); + ReaderSupplier readerSupplier = ReaderSupplierFactory.open(graphPath); + OnDiskGraphIndex diskGraph = OnDiskGraphIndex.load(readerSupplier); + + // ── Step 6: Search with random int8 query vectors ──────────────────────────── + // The score function reads each candidate's byte vector directly from disk — + // true int8 end-to-end, no float conversion anywhere in the search path. + int numQueries = 10; + int topK = 5; + Random rng = new Random(42); + + System.out.printf("%nSearching with %d random int8 query vectors (top-%d):%n", numQueries, topK); + + try (GraphSearcher searcher = new GraphSearcher(diskGraph)) { + OnDiskGraphIndex.View view = (OnDiskGraphIndex.View) searcher.getView(); + + for (int q = 0; q < numQueries; q++) { + ByteSequence queryBytes = randomInt8Vector(dimension, rng); + + // byteVectorRerankerFor reads each candidate's int8 vector from disk + // and scores it byte×byte — no float32 anywhere in the hot path. + SearchScoreProvider ssp = new DefaultSearchScoreProvider( + view.byteVectorRerankerFor(queryBytes, ByteVectorSimilarityFunction.EUCLIDEAN)); + + SearchResult result = searcher.search(ssp, topK, Bits.ALL); + + System.out.printf("Query %2d → top neighbors: ", q); + for (SearchResult.NodeScore ns : result.getNodes()) { + System.out.printf("(id=%d, score=%.4f) ", ns.node, ns.score); + } + System.out.println(); + } + } + + // cleanup + readerSupplier.close(); + Files.deleteIfExists(graphPath); + } + + /** + * Returns a random signed-byte vector as a {@link ByteSequence}. + * Each component is independently and uniformly drawn from [-128, 127]. + */ + private static ByteSequence randomInt8Vector(int dimension, Random rng) { + ByteSequence v = vts.createByteSequence(dimension); + for (int i = 0; i < dimension; i++) { + // nextInt(256) gives [0, 255]; subtract 128 → [-128, 127] + v.set(i, (byte) (rng.nextInt(256) - 128)); + } + return v; + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/TutorialRunner.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/TutorialRunner.java index c675f90c8..0e914e65d 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/TutorialRunner.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/TutorialRunner.java @@ -41,6 +41,9 @@ public static void main(String[] args) throws IOException { case "nvq": NvqExample.main(forwardArgs); break; + case "int8": + Int8Example.main(forwardArgs); + break; default: throw new IllegalArgumentException("Unknown example" + args[0]); } From d67617c8007737677e8665aea84cf3637b48236d Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Wed, 2 Sep 2026 08:59:26 +0000 Subject: [PATCH 09/13] Add int8 test coverage for byte-vector similarity, RABVV, BSP, SIMD, graph build, and disk round-trip --- .../TestListRandomAccessByteVectorValues.java | 83 +++++++ .../jvector/graph/TestVectorGraph.java | 162 +++++++++++++ .../graph/disk/TestOnDiskGraphIndex.java | 220 ++++++++++++++++++ .../similarity/BuildScoreProviderTest.java | 100 +++++++- .../TestByteVectorSimilarityFunction.java | 177 ++++++++++++++ .../vector/TestVectorizationProvider.java | 82 +++++++ 6 files changed, 823 insertions(+), 1 deletion(-) create mode 100644 jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestListRandomAccessByteVectorValues.java create mode 100644 jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestListRandomAccessByteVectorValues.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestListRandomAccessByteVectorValues.java new file mode 100644 index 000000000..81cab7fee --- /dev/null +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestListRandomAccessByteVectorValues.java @@ -0,0 +1,83 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.graph; + +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import org.junit.Test; + +import java.util.ArrayList; +import java.util.List; + +import static org.junit.Assert.*; + +public class TestListRandomAccessByteVectorValues { + + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + private ByteSequence seq(byte... values) { + return vts.createByteSequence(values); + } + + @Test + public void testSizeAndDimension() { + int dim = 4; + List> vectors = List.of(seq((byte) 1, (byte) 2, (byte) 3, (byte) 4), + seq((byte) 5, (byte) 6, (byte) 7, (byte) 8)); + var rabvv = new ListRandomAccessByteVectorValues(vectors, dim); + assertEquals(2, rabvv.size()); + assertEquals(dim, rabvv.dimension()); + } + + @Test + public void testGetVectorReturnsCorrectEntry() { + var v0 = seq((byte) 10, (byte) 20); + var v1 = seq((byte) -1, (byte) -2); + var v2 = seq((byte) 127, (byte) -128); + var rabvv = new ListRandomAccessByteVectorValues(List.of(v0, v1, v2), 2); + + assertSame(v0, rabvv.getVector(0)); + assertSame(v1, rabvv.getVector(1)); + assertSame(v2, rabvv.getVector(2)); + } + + @Test + public void testIsValueShared() { + var rabvv = new ListRandomAccessByteVectorValues(List.of(seq((byte) 0)), 1); + assertFalse(rabvv.isValueShared()); + } + + @Test + public void testCopyReturnsSelf() { + var rabvv = new ListRandomAccessByteVectorValues(List.of(seq((byte) 1, (byte) 2)), 2); + assertSame(rabvv, rabvv.copy()); + } + + @Test + public void testMutableBackingListIsReflected() { + // ListRandomAccessByteVectorValues documents that additions to the backing list are visible + List> backing = new ArrayList<>(); + backing.add(seq((byte) 1)); + var rabvv = new ListRandomAccessByteVectorValues(backing, 1); + assertEquals(1, rabvv.size()); + + backing.add(seq((byte) 2)); + assertEquals(2, rabvv.size()); + assertSame(backing.get(1), rabvv.getVector(1)); + } +} diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java index b784cd7d4..4f80eb24b 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestVectorGraph.java @@ -33,10 +33,13 @@ import io.github.jbellis.jvector.util.Bits; import io.github.jbellis.jvector.util.BoundedLongHeap; import io.github.jbellis.jvector.util.FixedBitSet; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import io.github.jbellis.jvector.graph.VectorValues; import org.junit.Before; import org.junit.Test; @@ -800,4 +803,163 @@ protected static Bits createRandomAcceptOrds(int startIndex, int length) { } return bits; } + + // ----------------------------------------------------------------------- + // Byte-vector (int8) graph construction tests + // ----------------------------------------------------------------------- + + /** Build a small int8 graph using every similarity function and verify it is navigable. */ + @Test + public void testByteVectorBuildAllSimilarityFunctions() { + for (var bvsf : ByteVectorSimilarityFunction.values()) { + testByteVectorBuild(bvsf, false); + testByteVectorBuild(bvsf, true); + } + } + + private void testByteVectorBuild(ByteVectorSimilarityFunction bvsf, boolean addHierarchy) { + int n = 50, dim = 8; + var rabvv = randomByteVectorValues(n, dim); + var builder = GraphIndexBuilder.builder(rabvv, bvsf, 8) + .withBeamWidth(20) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + var graph = builder.build(rabvv); + validateIndex(graph); + assertNotNull("entry node must be set", graph.getView().entryNode()); + assertEquals(n, graph.size(0)); + } + + /** addGraphNode(int, ByteSequence) inserts nodes one by one and yields the right count. */ + @Test + public void testByteVectorAddGraphNodeOneByOne() { + int n = 30, dim = 4; + var rabvv = randomByteVectorValues(n, dim); + var builder = GraphIndexBuilder.builder(rabvv, ByteVectorSimilarityFunction.EUCLIDEAN, 4) + .withBeamWidth(10) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + for (int i = 0; i < n; i++) { + builder.addGraphNode(i, rabvv.getVector(i)); + } + builder.cleanup(); + var graph = builder.getGraph(); + validateIndex(graph); + assertEquals(n, graph.size(0)); + // every node's neighbor count is within the declared degree + var view = graph.getView(); + for (int i = 0; i < n; i++) { + assertTrue(view.getNeighborsIterator(0, i).size() <= graph.getDegree(0)); + } + } + + /** + * addGraphNode(int, ByteSequence) must throw when the builder was constructed with a + * float-vector score provider. + */ + @Test + public void testByteAddGraphNodeThrowsOnFloatBuilder() { + var floatRavv = circularVectorValues(10); + var builder = new GraphIndexBuilder(floatRavv, VectorSimilarityFunction.EUCLIDEAN, 4, 10, 1.2f, 1.2f, false); + var bs = vectorTypeSupport.createByteSequence(4); + assertThrows(UnsupportedOperationException.class, () -> builder.addGraphNode(0, bs)); + } + + /** + * build(VectorValues) dispatches correctly when given a RandomAccessByteVectorValues. + */ + @Test + public void testByteVectorBuildViaGenericBuildMethod() { + int n = 40, dim = 6; + var rabvv = randomByteVectorValues(n, dim); + var builder = GraphIndexBuilder.builder(rabvv, ByteVectorSimilarityFunction.DOT_PRODUCT, 6) + .withBeamWidth(20) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + // call the generic VectorValues overload explicitly + var graph = builder.build((VectorValues) rabvv); + validateIndex(graph); + assertEquals(n, graph.size(0)); + } + + /** refineFinalGraph=false and refineFinalGraph=true both produce valid graphs. */ + @Test + public void testByteVectorRefineFinalGraphVariants() { + int n = 40, dim = 4; + for (boolean refine : new boolean[]{false, true}) { + var rabvv = randomByteVectorValues(n, dim); + // Use the 7-arg constructor (refineFinalGraph defaults to true); test both via the + // BSP constructor which exposes the refineFinalGraph knob + var bsp = io.github.jbellis.jvector.graph.similarity.BuildScoreProvider.byteVectorScoreProvider( + rabvv, ByteVectorSimilarityFunction.EUCLIDEAN); + var builder = new GraphIndexBuilder(bsp, dim, 4, 20, 1.2f, 1.2f, false, refine); + var graph = builder.build(rabvv); + validateIndex(graph); + assertEquals(n, graph.size(0)); + var view = graph.getView(); + int ub = graph.getIdUpperBound(); + for (int i = 0; i < n; i++) { + for (var it = view.getNeighborsIterator(0, i); it.hasNext(); ) { + int nb = it.nextInt(); + assertTrue("neighbor " + nb + " out of bounds", nb >= 0 && nb < ub); + } + } + } + } + + /** Smoke-test: a built byte-vector graph achieves reasonable top-5 recall. */ + @Test + public void testByteVectorSearchRecall() { + int n = 200, dim = 16; + int topK = 5; + var rabvv = randomByteVectorValues(n, dim); + var bvsf = ByteVectorSimilarityFunction.EUCLIDEAN; + var builder = GraphIndexBuilder.builder(rabvv, bvsf, 16) + .withBeamWidth(50) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + var graph = builder.build(rabvv); + + int totalMatches = 0; + int trials = 100; + for (int t = 0; t < trials; t++) { + byte[] rawQ = new byte[dim]; + getRandom().nextBytes(rawQ); + ByteSequence query = vectorTypeSupport.createByteSequence(rawQ); + + // brute-force top-K + NodeQueue expected = new NodeQueue(new BoundedLongHeap(topK), NodeQueue.Order.MIN_HEAP); + for (int i = 0; i < n; i++) { + expected.push(i, bvsf.compare(query, rabvv.getVector(i))); + } + + // graph search with efSearch = 100 (larger than topK to explore more candidates) + var ssp = new DefaultSearchScoreProvider( + (io.github.jbellis.jvector.graph.similarity.ScoreFunction.ExactScoreFunction) + node -> bvsf.compare(query, rabvv.getVector(node))); + var result = new GraphSearcher(graph).search(ssp, topK, 100, 0f, 0f, Bits.ALL); + var actualNodeIds = Arrays.stream(result.getNodes(), 0, topK) + .mapToInt(ns -> ns.node).toArray(); + + totalMatches += computeOverlap(actualNodeIds, expected.nodesCopy()); + } + // with efSearch=100 over a 200-node graph, we can visit most nodes; expect >=90% recall + double overlap = totalMatches / (double) (trials * topK); + assertTrue("byte-vector recall " + overlap + " is too low", overlap > 0.9); + } + + /** Creates a list-backed RandomAccessByteVectorValues with random int8 vectors. */ + private ListRandomAccessByteVectorValues randomByteVectorValues(int n, int dim) { + var list = new java.util.ArrayList>(n); + for (int i = 0; i < n; i++) { + byte[] raw = new byte[dim]; + getRandom().nextBytes(raw); + list.add(vectorTypeSupport.createByteSequence(raw)); + } + return new ListRandomAccessByteVectorValues(list, dim); + } } diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java index c84d5b103..47328e38a 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/disk/TestOnDiskGraphIndex.java @@ -23,6 +23,7 @@ import io.github.jbellis.jvector.graph.GraphIndexBuilder; import io.github.jbellis.jvector.graph.GraphSearcher; import io.github.jbellis.jvector.graph.ImmutableGraphIndex; +import io.github.jbellis.jvector.graph.ListRandomAccessByteVectorValues; import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; import io.github.jbellis.jvector.graph.NodesIterator; import io.github.jbellis.jvector.graph.RandomAccessVectorValues; @@ -30,6 +31,7 @@ import io.github.jbellis.jvector.graph.disk.feature.Feature; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; import io.github.jbellis.jvector.graph.disk.feature.FusedPQ; +import io.github.jbellis.jvector.graph.disk.feature.InlineByteVectors; import io.github.jbellis.jvector.graph.disk.feature.InlineVectors; import io.github.jbellis.jvector.graph.disk.feature.NVQ; import io.github.jbellis.jvector.graph.disk.feature.SeparatedNVQ; @@ -38,7 +40,13 @@ import io.github.jbellis.jvector.quantization.PQVectors; import io.github.jbellis.jvector.quantization.ProductQuantization; import io.github.jbellis.jvector.util.Bits; +import io.github.jbellis.jvector.graph.similarity.DefaultSearchScoreProvider; +import io.github.jbellis.jvector.graph.similarity.ScoreFunction; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; import org.junit.After; import org.junit.Before; import org.junit.Test; @@ -572,4 +580,216 @@ public void testIncrementalWrites() throws IOException { throw new RuntimeException(e); } } + + // ----------------------------------------------------------------------- + // InlineByteVectors (int8) tests + // ----------------------------------------------------------------------- + + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + /** Builds a list of n random int8 vectors of the given dimension. */ + private List> randomByteVectors(int n, int dim) { + var list = new ArrayList>(n); + for (int i = 0; i < n; i++) { + byte[] raw = new byte[dim]; + getRandom().nextBytes(raw); + list.add(vts.createByteSequence(raw)); + } + return list; + } + + /** + * Full round-trip: build with InlineByteVectors, write, reload, verify + * getByteVector() returns the original bytes for every node. + */ + @Test + public void testInlineByteVectorsRoundTrip() throws IOException { + int n = 50, dim = 8; + var vectors = randomByteVectors(n, dim); + var rabvv = new ListRandomAccessByteVectorValues(vectors, dim); + + // Build a graph + var builder = GraphIndexBuilder.builder(rabvv, ByteVectorSimilarityFunction.EUCLIDEAN, 8) + .withBeamWidth(20) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + var graph = builder.build(rabvv); + + // Write with InlineByteVectors feature + var outputPath = testDirectory.resolve("byte_vectors_roundtrip"); + try (var writer = new OnDiskGraphIndexWriter.Builder(graph, outputPath) + .with(new InlineByteVectors(dim)) + .build()) + { + writer.write(Feature.singleStateFactory(FeatureId.INLINE_BYTE_VECTORS, + nodeId -> new InlineByteVectors.State(rabvv.getVector(nodeId)))); + } + + // Reload and verify every vector + try (var rs = new SimpleMappedReader.Supplier(outputPath.toAbsolutePath()); + var onDisk = OnDiskGraphIndex.load(rs); + var view = onDisk.getView()) + { + assertTrue("INLINE_BYTE_VECTORS feature missing", + onDisk.getFeatureSet().contains(FeatureId.INLINE_BYTE_VECTORS)); + assertEquals(dim, onDisk.dimension); + + for (int i = 0; i < n; i++) { + var expected = rabvv.getVector(i); + var actual = view.getByteVector(i); + assertEquals("byte vector length mismatch at node " + i, expected.length(), actual.length()); + for (int d = 0; d < dim; d++) { + assertEquals("byte mismatch at node " + i + " dim " + d, + expected.get(d), actual.get(d)); + } + } + } + } + + /** + * featureSize() of InlineByteVectors must equal the dimension (1 byte per component), + * confirming 4× compression vs float32 InlineVectors. + */ + @Test + public void testInlineByteVectorsFeatureSize() { + for (int dim : new int[]{1, 8, 128, 256}) { + var ibv = new InlineByteVectors(dim); + assertEquals("featureSize should equal dimension", dim, ibv.featureSize()); + // float32 InlineVectors uses 4 * dim bytes + assertEquals("byte feature 4x smaller than float32", + new InlineVectors(dim).featureSize(), 4 * ibv.featureSize()); + } + } + + /** + * getByteVector throws UnsupportedOperationException when the graph has no INLINE_BYTE_VECTORS feature. + */ + @Test + public void testGetByteVectorThrowsWithoutFeature() throws IOException { + // Build and write a float-vector graph (no INLINE_BYTE_VECTORS) + var graph = new TestUtil.RandomlyConnectedGraphIndex(10, 4, getRandom()); + var ravv = new TestVectorGraph.CircularFloatVectorValues(10); + var outputPath = testDirectory.resolve("float_graph_no_bytes"); + TestUtil.writeGraph(graph, ravv, outputPath); + + try (var rs = new SimpleMappedReader.Supplier(outputPath.toAbsolutePath()); + var onDisk = OnDiskGraphIndex.load(rs); + var view = onDisk.getView()) + { + assertFalse(onDisk.getFeatureSet().contains(FeatureId.INLINE_BYTE_VECTORS)); + assertThrows(UnsupportedOperationException.class, () -> view.getByteVector(0)); + } + } + + /** + * byteVectorRerankerFor throws UnsupportedOperationException when the graph has no INLINE_BYTE_VECTORS feature. + */ + @Test + public void testByteVectorRerankerThrowsWithoutFeature() throws IOException { + var graph = new TestUtil.RandomlyConnectedGraphIndex(10, 4, getRandom()); + var ravv = new TestVectorGraph.CircularFloatVectorValues(10); + var outputPath = testDirectory.resolve("float_graph_reranker"); + TestUtil.writeGraph(graph, ravv, outputPath); + + try (var rs = new SimpleMappedReader.Supplier(outputPath.toAbsolutePath()); + var onDisk = OnDiskGraphIndex.load(rs); + var view = onDisk.getView()) + { + var qb = vts.createByteSequence(2); + assertThrows(UnsupportedOperationException.class, + () -> view.byteVectorRerankerFor(qb, ByteVectorSimilarityFunction.EUCLIDEAN)); + } + } + + /** + * byteVectorRerankerFor returns scores consistent with a direct bvsf.compare() call. + */ + @Test + public void testByteVectorRerankerScoresMatchDirectCompare() throws IOException { + int n = 30, dim = 4; + var vectors = randomByteVectors(n, dim); + var rabvv = new ListRandomAccessByteVectorValues(vectors, dim); + var bvsf = ByteVectorSimilarityFunction.EUCLIDEAN; + + var builder = GraphIndexBuilder.builder(rabvv, bvsf, 8) + .withBeamWidth(20) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + var graph = builder.build(rabvv); + + var outputPath = testDirectory.resolve("byte_reranker_test"); + try (var writer = new OnDiskGraphIndexWriter.Builder(graph, outputPath) + .with(new InlineByteVectors(dim)) + .build()) + { + writer.write(Feature.singleStateFactory(FeatureId.INLINE_BYTE_VECTORS, + nodeId -> new InlineByteVectors.State(rabvv.getVector(nodeId)))); + } + + try (var rs = new SimpleMappedReader.Supplier(outputPath.toAbsolutePath()); + var onDisk = OnDiskGraphIndex.load(rs); + var view = onDisk.getView()) + { + byte[] rawQ = new byte[dim]; + getRandom().nextBytes(rawQ); + var query = vts.createByteSequence(rawQ); + var reranker = view.byteVectorRerankerFor(query, bvsf); + + for (int i = 0; i < n; i++) { + float expected = bvsf.compare(query, rabvv.getVector(i)); + float actual = reranker.similarityTo(i); + assertEquals("reranker score mismatch at node " + i, expected, actual, 1e-5f); + } + } + } + + /** + * All three byte-vector similarity functions produce results in [0, 1] for all nodes + * after a full int8 end-to-end write/search cycle. + */ + @Test + public void testByteVectorEndToEndAllSimilarityFunctions() throws IOException { + int n = 40, dim = 8; + + for (var bvsf : ByteVectorSimilarityFunction.values()) { + var vectors = randomByteVectors(n, dim); + var rabvv = new ListRandomAccessByteVectorValues(vectors, dim); + + var builder = GraphIndexBuilder.builder(rabvv, bvsf, 8) + .withBeamWidth(20) + .withNeighborOverflow(1.2f) + .withAlpha(1.2f) + .build(); + var graph = builder.build(rabvv); + + var outputPath = testDirectory.resolve("byte_e2e_" + bvsf.name()); + try (var writer = new OnDiskGraphIndexWriter.Builder(graph, outputPath) + .with(new InlineByteVectors(dim)) + .build()) + { + writer.write(Feature.singleStateFactory(FeatureId.INLINE_BYTE_VECTORS, + nodeId -> new InlineByteVectors.State(rabvv.getVector(nodeId)))); + } + + try (var rs = new SimpleMappedReader.Supplier(outputPath.toAbsolutePath()); + var onDisk = OnDiskGraphIndex.load(rs); + var view = onDisk.getView()) + { + byte[] rawQ = new byte[dim]; + getRandom().nextBytes(rawQ); + var query = vts.createByteSequence(rawQ); + var reranker = view.byteVectorRerankerFor(query, bvsf); + var ssp = new DefaultSearchScoreProvider(reranker); + var result = new GraphSearcher(onDisk).search(ssp, 5, Bits.ALL); + + for (var ns : result.getNodes()) { + float score = ns.score; + assertTrue(bvsf + " score out of [0,1] for node " + ns.node + ": " + score, + score >= 0f && score <= 1.0f); + } + } + } + } } diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProviderTest.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProviderTest.java index 4942b8efb..542a71b45 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProviderTest.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProviderTest.java @@ -16,9 +16,12 @@ package io.github.jbellis.jvector.graph.similarity; +import io.github.jbellis.jvector.graph.ListRandomAccessByteVectorValues; import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; import org.junit.Test; @@ -26,7 +29,8 @@ import java.util.ArrayList; import java.util.List; -import static org.junit.Assert.assertEquals; +import static org.junit.Assert.*; +import static org.junit.Assert.assertThrows; public class BuildScoreProviderTest { private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); @@ -69,4 +73,98 @@ public void testOrdinalMapping() { var dsp0 = bsp.diversityProviderFor(0); assertEquals(vsf.compare(vectors.get(2), vectors.get(0)), dsp0.exactScoreFunction().similarityTo(1), 1e-6f); } + + // ----------------------------------------------------------------------- + // byteVectorScoreProvider tests + // ----------------------------------------------------------------------- + + private ByteSequence bseq(byte... values) { + return vts.createByteSequence(values); + } + + /** Helper: build a small RABVV with three 2-D signed-byte vectors. */ + private ListRandomAccessByteVectorValues byteRavv() { + return new ListRandomAccessByteVectorValues( + List.of(bseq((byte) 10, (byte) 0), // node 0 + bseq((byte) 0, (byte) 10), // node 1 + bseq((byte) -10, (byte) 0)), // node 2 + 2); + } + + @Test + public void testByteVectorIsExact() { + var bsp = BuildScoreProvider.byteVectorScoreProvider(byteRavv(), ByteVectorSimilarityFunction.EUCLIDEAN); + assertTrue(bsp.isExact()); + } + + @Test + public void testByteVectorSearchProviderForByteSequence() { + var bvsf = ByteVectorSimilarityFunction.EUCLIDEAN; + var rabvv = byteRavv(); + var bsp = BuildScoreProvider.byteVectorScoreProvider(rabvv, bvsf); + + // searchProviderFor(ByteSequence) scores all nodes against the given query + ByteSequence query = bseq((byte) 10, (byte) 0); // identical to node 0 + var ssp = bsp.searchProviderFor(query); + + // node 0 should be self-similar (score = 1.0 for EUCLIDEAN with zero distance) + assertEquals(1.0f, ssp.exactScoreFunction().similarityTo(0), 1e-5f); + // node 2 = [-10, 0], distance vs [10,0] = 400, not 1.0 + assertTrue(ssp.exactScoreFunction().similarityTo(2) < 1.0f); + } + + @Test + public void testByteVectorSearchProviderForNode() { + var bvsf = ByteVectorSimilarityFunction.DOT_PRODUCT; + var rabvv = byteRavv(); + var bsp = BuildScoreProvider.byteVectorScoreProvider(rabvv, bvsf); + + // searchProviderFor(int) should delegate to searchProviderFor(ByteSequence) + var sspByNode = bsp.searchProviderFor(0); + ByteSequence v0 = rabvv.getVector(0); + var sspBySeq = bsp.searchProviderFor(v0); + assertEquals(sspByNode.exactScoreFunction().similarityTo(1), + sspBySeq.exactScoreFunction().similarityTo(1), 1e-6f); + } + + @Test + public void testByteVectorDiversityProviderMatchesSearch() { + var bvsf = ByteVectorSimilarityFunction.COSINE; + var bsp = BuildScoreProvider.byteVectorScoreProvider(byteRavv(), bvsf); + + // diversityProviderFor delegates to searchProviderFor(int) + var search = bsp.searchProviderFor(1); + var diversity = bsp.diversityProviderFor(1); + assertEquals(search.exactScoreFunction().similarityTo(0), + diversity.exactScoreFunction().similarityTo(0), 1e-6f); + } + + @Test + public void testByteVectorDiversityScoreFunction() { + var bvsf = ByteVectorSimilarityFunction.EUCLIDEAN; + var rabvv = byteRavv(); + var bsp = BuildScoreProvider.byteVectorScoreProvider(rabvv, bvsf); + + // diversityScoreFunctionFor(n1).similarityTo(n2) == bvsf.compare(v_n1, v_n2) + var dsf = bsp.diversityScoreFunctionFor(0); + assertEquals(bvsf.compare(rabvv.getVector(0), rabvv.getVector(2)), + dsf.similarityTo(2), 1e-6f); + } + + @Test + public void testByteVectorApproximateCentroid() { + // centroid of [10,0], [0,10], [-10,0] should be [0, 10/3] + var bsp = BuildScoreProvider.byteVectorScoreProvider(byteRavv(), ByteVectorSimilarityFunction.EUCLIDEAN); + var centroid = bsp.approximateCentroid(); + assertEquals(2, centroid.length()); + assertEquals(0.0f, centroid.get(0), 1e-5f); + assertEquals(10.0f / 3.0f, centroid.get(1), 1e-5f); + } + + @Test + public void testByteVectorThrowsForFloatQuery() { + var bsp = BuildScoreProvider.byteVectorScoreProvider(byteRavv(), ByteVectorSimilarityFunction.EUCLIDEAN); + VectorFloat floatQuery = vts.createFloatVector(new float[]{1.0f, 0.0f}); + assertThrows(UnsupportedOperationException.class, () -> bsp.searchProviderFor(floatQuery)); + } } \ No newline at end of file diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java new file mode 100644 index 000000000..c95623b69 --- /dev/null +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java @@ -0,0 +1,177 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.vector; + +import com.carrotsearch.randomizedtesting.RandomizedTest; +import io.github.jbellis.jvector.vector.types.ByteSequence; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import org.junit.Test; + +import static org.junit.Assert.*; + +public class TestByteVectorSimilarityFunction extends RandomizedTest { + + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + private ByteSequence seq(byte... values) { + return vts.createByteSequence(values); + } + + // ----------------------------------------------------------------------- + // EUCLIDEAN + // ----------------------------------------------------------------------- + + @Test + public void testEuclideanKnownValue() { + // v1=[1,0], v2=[0,1] squared-L2 = 1+1 = 2 + // maxSquaredDist = 2 * 255^2 = 130050 + // expected = 1 / (1 + 2/130050) + var v1 = seq((byte) 1, (byte) 0); + var v2 = seq((byte) 0, (byte) 1); + float squaredL2 = 2.0f; + float max = 2 * 255.0f * 255.0f; + float expected = 1.0f / (1.0f + squaredL2 / max); + assertEquals(expected, ByteVectorSimilarityFunction.EUCLIDEAN.compare(v1, v2), 1e-6f); + } + + @Test + public void testEuclideanIdenticalVectors() { + var v = seq((byte) 42, (byte) -7, (byte) 100); + // squaredL2 = 0, so score = 1/(1+0) = 1.0 + assertEquals(1.0f, ByteVectorSimilarityFunction.EUCLIDEAN.compare(v, v), 1e-6f); + } + + @Test + public void testEuclideanResultInRange() { + for (int trial = 0; trial < 50; trial++) { + byte[] raw1 = new byte[32]; + byte[] raw2 = new byte[32]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + float score = ByteVectorSimilarityFunction.EUCLIDEAN.compare(seq(raw1), seq(raw2)); + assertTrue("EUCLIDEAN score out of (0,1]: " + score, score > 0f && score <= 1.0f); + } + } + + @Test + public void testEuclideanSymmetry() { + byte[] raw1 = new byte[16]; + byte[] raw2 = new byte[16]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + assertEquals( + ByteVectorSimilarityFunction.EUCLIDEAN.compare(seq(raw1), seq(raw2)), + ByteVectorSimilarityFunction.EUCLIDEAN.compare(seq(raw2), seq(raw1)), + 1e-6f); + } + + // ----------------------------------------------------------------------- + // DOT_PRODUCT + // ----------------------------------------------------------------------- + + @Test + public void testDotProductKnownValue() { + // v1=[1,0], v2=[0,1] dot = 0 + // maxMag = 2 * 127^2 = 32258 + // expected = (1 + 0/32258) / 2 = 0.5 + var v1 = seq((byte) 1, (byte) 0); + var v2 = seq((byte) 0, (byte) 1); + assertEquals(0.5f, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(v1, v2), 1e-6f); + } + + @Test + public void testDotProductResultInRange() { + for (int trial = 0; trial < 50; trial++) { + byte[] raw1 = new byte[32]; + byte[] raw2 = new byte[32]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + float score = ByteVectorSimilarityFunction.DOT_PRODUCT.compare(seq(raw1), seq(raw2)); + assertTrue("DOT_PRODUCT score out of [0,1]: " + score, score >= 0f && score <= 1.0f); + } + } + + @Test + public void testDotProductSymmetry() { + byte[] raw1 = new byte[16]; + byte[] raw2 = new byte[16]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + assertEquals( + ByteVectorSimilarityFunction.DOT_PRODUCT.compare(seq(raw1), seq(raw2)), + ByteVectorSimilarityFunction.DOT_PRODUCT.compare(seq(raw2), seq(raw1)), + 1e-6f); + } + + // ----------------------------------------------------------------------- + // COSINE + // ----------------------------------------------------------------------- + + @Test + public void testCosineParallelVectors() { + // v1 == v2 → cosine = 1.0, score = (1+1)/2 = 1.0 + var v = seq((byte) 3, (byte) 4); + assertEquals(1.0f, ByteVectorSimilarityFunction.COSINE.compare(v, v), 1e-5f); + } + + @Test + public void testCosineOrthogonalVectors() { + // [1,0] · [0,1] = 0, cosine = 0, score = (1+0)/2 = 0.5 + var v1 = seq((byte) 1, (byte) 0); + var v2 = seq((byte) 0, (byte) 1); + assertEquals(0.5f, ByteVectorSimilarityFunction.COSINE.compare(v1, v2), 1e-5f); + } + + @Test + public void testCosineResultInRange() { + for (int trial = 0; trial < 50; trial++) { + byte[] raw1 = new byte[32]; + byte[] raw2 = new byte[32]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + float score = ByteVectorSimilarityFunction.COSINE.compare(seq(raw1), seq(raw2)); + assertTrue("COSINE score out of [0,1]: " + score, score >= 0f && score <= 1.0f); + } + } + + @Test + public void testCosineSymmetry() { + byte[] raw1 = new byte[16]; + byte[] raw2 = new byte[16]; + getRandom().nextBytes(raw1); + getRandom().nextBytes(raw2); + assertEquals( + ByteVectorSimilarityFunction.COSINE.compare(seq(raw1), seq(raw2)), + ByteVectorSimilarityFunction.COSINE.compare(seq(raw2), seq(raw1)), + 1e-5f); + } + + // ----------------------------------------------------------------------- + // Boundary values + // ----------------------------------------------------------------------- + + @Test + public void testAllMaxValues() { + // All 127 vectors — EUCLIDEAN identity = 1, DOT_PRODUCT = 1, COSINE = 1 + byte[] raw = new byte[8]; + java.util.Arrays.fill(raw, (byte) 127); + var v = seq(raw); + assertEquals(1.0f, ByteVectorSimilarityFunction.EUCLIDEAN.compare(v, v), 1e-5f); + assertEquals(1.0f, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(v, v), 1e-5f); + assertEquals(1.0f, ByteVectorSimilarityFunction.COSINE.compare(v, v), 1e-5f); + } +} diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java index 29f0ec18f..10e61f3fe 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestVectorizationProvider.java @@ -94,6 +94,88 @@ public void testSimilarityMetricsByte() { 0.0001f); } + /** + * Verifies that the SIMD byte-vector kernels agree with the scalar baseline for + * several dimensions that stress boundary conditions in the vectorised loops: + *

+ */ + @Test + public void testSimilarityMetricsByteEdgeDimensions() { + Assume.assumeTrue(hasSimd); + + VectorizationProvider scalar = new DefaultVectorizationProvider(); + VectorizationProvider simd = VectorizationProvider.getInstance(); + + for (int dim : new int[]{1, 8, 128, 255, 256}) { + byte[] rawA = new byte[dim]; + byte[] rawB = new byte[dim]; + getRandom().nextBytes(rawA); + getRandom().nextBytes(rawB); + + ByteSequence sA = scalar.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence sB = scalar.getVectorTypeSupport().createByteSequence(rawB); + ByteSequence vA = simd.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence vB = simd.getVectorTypeSupport().createByteSequence(rawB); + + Assert.assertEquals("dotProduct dim=" + dim, + scalar.getVectorUtilSupport().dotProduct(sA, sB), + simd.getVectorUtilSupport().dotProduct(vA, vB), + 0.0001f); + Assert.assertEquals("squareDistance dim=" + dim, + scalar.getVectorUtilSupport().squareDistance(sA, sB), + simd.getVectorUtilSupport().squareDistance(vA, vB), + 0.0001f); + Assert.assertEquals("cosine dim=" + dim, + scalar.getVectorUtilSupport().cosine(sA, sB), + simd.getVectorUtilSupport().cosine(vA, vB), + 0.0001f); + } + } + + /** + * SIMD kernels must treat bytes as signed. This test uses vectors whose + * correct result depends on negative byte values to catch sign-extension bugs. + */ + @Test + public void testSimilarityMetricsByteSignHandling() { + Assume.assumeTrue(hasSimd); + + VectorizationProvider scalar = new DefaultVectorizationProvider(); + VectorizationProvider simd = VectorizationProvider.getInstance(); + + // mix of extreme signed values: max positive 127, min negative -128 + byte[] rawA = new byte[16]; + byte[] rawB = new byte[16]; + for (int i = 0; i < 16; i++) { + rawA[i] = (i % 2 == 0) ? (byte) 127 : (byte) -128; + rawB[i] = (i % 2 == 0) ? (byte) -128 : (byte) 127; + } + + ByteSequence sA = scalar.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence sB = scalar.getVectorTypeSupport().createByteSequence(rawB); + ByteSequence vA = simd.getVectorTypeSupport().createByteSequence(rawA); + ByteSequence vB = simd.getVectorTypeSupport().createByteSequence(rawB); + + Assert.assertEquals("signed dotProduct", + scalar.getVectorUtilSupport().dotProduct(sA, sB), + simd.getVectorUtilSupport().dotProduct(vA, vB), + 0.0001f); + Assert.assertEquals("signed squareDistance", + scalar.getVectorUtilSupport().squareDistance(sA, sB), + simd.getVectorUtilSupport().squareDistance(vA, vB), + 0.0001f); + Assert.assertEquals("signed cosine", + scalar.getVectorUtilSupport().cosine(sA, sB), + simd.getVectorUtilSupport().cosine(vA, vB), + 0.0001f); + } + @Test public void testAssembleAndSum() { Assume.assumeTrue(hasSimd); From d0e9cba780a31a9a39d29d67d6e92d38895efe41 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Tue, 15 Sep 2026 07:53:30 +0000 Subject: [PATCH 10/13] Remove dead offset locals from byte similarity methods --- .../vector/PanamaVectorUtilSupport.java | 60 +++++++------------ 1 file changed, 21 insertions(+), 39 deletions(-) diff --git a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java index df49f8858..84f337187 100644 --- a/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java +++ b/jvector-twenty/src/main/java/io/github/jbellis/jvector/vector/PanamaVectorUtilSupport.java @@ -1002,17 +1002,15 @@ public float dotProduct(ByteSequence a, ByteSequence b) { float dotProductBytes512(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_128.length(); // 16 final int limit = ByteVector.SPECIES_128.loopBound(length); IntVector acc = IntVector.zero(IntVector.SPECIES_512); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); acc = acc.add(va.mul(vb)); @@ -1020,24 +1018,22 @@ float dotProductBytes512(ByteSequence a, ByteSequence b) { int result = acc.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - result += a.get(aOff + i) * b.get(bOff + i); + result += a.get(i) * b.get(i); } return result; } float dotProductBytes256(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_64.length(); // 8 final int limit = ByteVector.SPECIES_64.loopBound(length); IntVector acc = IntVector.zero(IntVector.SPECIES_256); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); acc = acc.add(va.mul(vb)); @@ -1045,7 +1041,7 @@ float dotProductBytes256(ByteSequence a, ByteSequence b) { int result = acc.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - result += a.get(aOff + i) * b.get(bOff + i); + result += a.get(i) * b.get(i); } return result; } @@ -1053,11 +1049,9 @@ float dotProductBytes256(ByteSequence a, ByteSequence b) { float dotProductBytes128(ByteSequence a, ByteSequence b) { // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); int result = 0; for (int i = 0; i < length; i++) { - result += a.get(aOff + i) * b.get(bOff + i); + result += a.get(i) * b.get(i); } return result; } @@ -1076,17 +1070,15 @@ public float squareDistance(ByteSequence a, ByteSequence b) { float squareDistanceBytes512(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_128.length(); final int limit = ByteVector.SPECIES_128.loopBound(length); IntVector acc = IntVector.zero(IntVector.SPECIES_512); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); IntVector diff = va.sub(vb); @@ -1095,7 +1087,7 @@ float squareDistanceBytes512(ByteSequence a, ByteSequence b) { int result = acc.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - int diff = a.get(aOff + i) - b.get(bOff + i); + int diff = a.get(i) - b.get(i); result += diff * diff; } return result; @@ -1103,17 +1095,15 @@ float squareDistanceBytes512(ByteSequence a, ByteSequence b) { float squareDistanceBytes256(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_64.length(); final int limit = ByteVector.SPECIES_64.loopBound(length); IntVector acc = IntVector.zero(IntVector.SPECIES_256); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); IntVector diff = va.sub(vb); @@ -1122,7 +1112,7 @@ float squareDistanceBytes256(ByteSequence a, ByteSequence b) { int result = acc.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - int diff = a.get(aOff + i) - b.get(bOff + i); + int diff = a.get(i) - b.get(i); result += diff * diff; } return result; @@ -1131,11 +1121,9 @@ float squareDistanceBytes256(ByteSequence a, ByteSequence b) { float squareDistanceBytes128(ByteSequence a, ByteSequence b) { // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); int result = 0; for (int i = 0; i < length; i++) { - int diff = a.get(aOff + i) - b.get(bOff + i); + int diff = a.get(i) - b.get(i); result += diff * diff; } return result; @@ -1155,8 +1143,6 @@ public float cosine(ByteSequence a, ByteSequence b) { float cosineBytes512(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_128.length(); final int limit = ByteVector.SPECIES_128.loopBound(length); IntVector dot = IntVector.zero(IntVector.SPECIES_512); @@ -1164,10 +1150,10 @@ float cosineBytes512(ByteSequence a, ByteSequence b) { IntVector normB = IntVector.zero(IntVector.SPECIES_512); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_128, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) .reinterpretAsInts(); dot = dot.add(va.mul(vb)); @@ -1180,7 +1166,7 @@ float cosineBytes512(ByteSequence a, ByteSequence b) { long normBResult = normB.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - int ai = a.get(aOff + i), bi = b.get(bOff + i); + int ai = a.get(i), bi = b.get(i); dotResult += ai * bi; normAResult += ai * ai; normBResult += bi * bi; @@ -1190,8 +1176,6 @@ float cosineBytes512(ByteSequence a, ByteSequence b) { float cosineBytes256(ByteSequence a, ByteSequence b) { final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); final int step = ByteVector.SPECIES_64.length(); final int limit = ByteVector.SPECIES_64.loopBound(length); IntVector dot = IntVector.zero(IntVector.SPECIES_256); @@ -1199,10 +1183,10 @@ float cosineBytes256(ByteSequence a, ByteSequence b) { IntVector normB = IntVector.zero(IntVector.SPECIES_256); for (int i = 0; i < limit; i += step) { - IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, aOff + i) + IntVector va = fromByteSequence(ByteVector.SPECIES_64, a, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); - IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, bOff + i) + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, i) .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) .reinterpretAsInts(); dot = dot.add(va.mul(vb)); @@ -1215,7 +1199,7 @@ float cosineBytes256(ByteSequence a, ByteSequence b) { long normBResult = normB.reduceLanes(VectorOperators.ADD); for (int i = limit; i < length; i++) { - int ai = a.get(aOff + i), bi = b.get(bOff + i); + int ai = a.get(i), bi = b.get(i); dotResult += ai * bi; normAResult += ai * ai; normBResult += bi * bi; @@ -1226,11 +1210,9 @@ float cosineBytes256(ByteSequence a, ByteSequence b) { float cosineBytes128(ByteSequence a, ByteSequence b) { // ByteVector.SPECIES_32 does not exist; scalar is fastest at 128-bit width final int length = a.length(); - final int aOff = a.offset(); - final int bOff = b.offset(); long dotResult = 0, normAResult = 0, normBResult = 0; for (int i = 0; i < length; i++) { - int ai = a.get(aOff + i), bi = b.get(bOff + i); + int ai = a.get(i), bi = b.get(i); dotResult += ai * bi; normAResult += ai * ai; normBResult += bi * bi; From 622980c708c9fa7ced2407c49e733ae3590c7232 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Tue, 15 Sep 2026 07:55:49 +0000 Subject: [PATCH 11/13] Fix maxMagnitude calculation for ByteVectorSimilarityFunction.DOT_PRODUCT --- .../vector/ByteVectorSimilarityFunction.java | 8 +++--- .../TestByteVectorSimilarityFunction.java | 27 ++++++++++++------- 2 files changed, 22 insertions(+), 13 deletions(-) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java index 2390343ea..33d2875ef 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/vector/ByteVectorSimilarityFunction.java @@ -43,15 +43,15 @@ public float compare(ByteSequence v1, ByteSequence v2) { /** * Dot product normalised to {@code [0, 1]}. - * Raw int8 dot product is divided by {@code n * 127^2} (the maximum possible magnitude) - * before applying the {@code (1 + x) / 2} mapping, so the result is always in [0, 1] - * regardless of dimension or whether the vectors are unit-norm. + * Raw int8 dot product is divided by {@code n * 128^2} (the maximum possible magnitude, + * achieved when components are -128) before applying the {@code (1 + x) / 2} mapping, + * so the result is always in [0, 1] regardless of dimension or whether the vectors are unit-norm. * For already unit-norm int8 vectors (e.g. Cohere, OpenAI reduced-precision) prefer {@link #COSINE}. */ DOT_PRODUCT { @Override public float compare(ByteSequence v1, ByteSequence v2) { - float maxMagnitude = v1.length() * (127.0f * 127.0f); + float maxMagnitude = v1.length() * (128.0f * 128.0f); return (1.0f + VectorUtil.dotProduct(v1, v2) / maxMagnitude) / 2.0f; } }, diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java index c95623b69..8a3e08e8f 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java @@ -86,8 +86,8 @@ public void testEuclideanSymmetry() { @Test public void testDotProductKnownValue() { // v1=[1,0], v2=[0,1] dot = 0 - // maxMag = 2 * 127^2 = 32258 - // expected = (1 + 0/32258) / 2 = 0.5 + // maxMag = 2 * 128^2 = 32768 + // expected = (1 + 0/32768) / 2 = 0.5 var v1 = seq((byte) 1, (byte) 0); var v2 = seq((byte) 0, (byte) 1); assertEquals(0.5f, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(v1, v2), 1e-6f); @@ -166,12 +166,21 @@ public void testCosineSymmetry() { @Test public void testAllMaxValues() { - // All 127 vectors — EUCLIDEAN identity = 1, DOT_PRODUCT = 1, COSINE = 1 - byte[] raw = new byte[8]; - java.util.Arrays.fill(raw, (byte) 127); - var v = seq(raw); - assertEquals(1.0f, ByteVectorSimilarityFunction.EUCLIDEAN.compare(v, v), 1e-5f); - assertEquals(1.0f, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(v, v), 1e-5f); - assertEquals(1.0f, ByteVectorSimilarityFunction.COSINE.compare(v, v), 1e-5f); + // All -128 vectors — EUCLIDEAN identity = 1, DOT_PRODUCT = 1, COSINE = 1 + byte[] rawMin = new byte[8]; + java.util.Arrays.fill(rawMin, (byte) -128); + var vMin = seq(rawMin); + assertEquals(1.0f, ByteVectorSimilarityFunction.EUCLIDEAN.compare(vMin, vMin), 1e-5f); + assertEquals(1.0f, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(vMin, vMin), 1e-5f); + assertEquals(1.0f, ByteVectorSimilarityFunction.COSINE.compare(vMin, vMin), 1e-5f); + + // All 127 vectors — EUCLIDEAN identity = 1, COSINE = 1 + byte[] rawMax = new byte[8]; + java.util.Arrays.fill(rawMax, (byte) 127); + var vMax = seq(rawMax); + assertEquals(1.0f, ByteVectorSimilarityFunction.EUCLIDEAN.compare(vMax, vMax), 1e-5f); + float expectedDot = (1.0f + (127.0f * 127.0f) / (128.0f * 128.0f)) / 2.0f; + assertEquals(expectedDot, ByteVectorSimilarityFunction.DOT_PRODUCT.compare(vMax, vMax), 1e-5f); + assertEquals(1.0f, ByteVectorSimilarityFunction.COSINE.compare(vMax, vMax), 1e-5f); } } From 3d14143eaae5aedf88095b29ff28d910dd91d126 Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Tue, 15 Sep 2026 07:55:54 +0000 Subject: [PATCH 12/13] Extract asRandomAccessSupplier with type check in BuildScoreProvider --- .../graph/similarity/BuildScoreProvider.java | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java index 634a6f8cf..2c7ffa8cd 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/similarity/BuildScoreProvider.java @@ -19,6 +19,7 @@ import io.github.jbellis.jvector.graph.RandomAccessByteVectorValues; import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import io.github.jbellis.jvector.graph.RemappedRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.VectorValues; import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.quantization.BQVectors; import io.github.jbellis.jvector.quantization.PQVectors; @@ -126,8 +127,8 @@ static BuildScoreProvider randomAccessScoreProvider(RandomAccessVectorValues rav // colliding. ThreadLocalSupplier makes this a no-op if the RAVV is actually un-shared. var vectorsRaw = ravv.threadLocalSupplier(); var vectorsCopyRaw = ravv.threadLocalSupplier(); - Supplier vectors = () -> (RandomAccessVectorValues) vectorsRaw.get(); - Supplier vectorsCopy = () -> (RandomAccessVectorValues) vectorsCopyRaw.get(); + Supplier vectors = asRandomAccessSupplier(vectorsRaw); + Supplier vectorsCopy = asRandomAccessSupplier(vectorsCopyRaw); return new BuildScoreProvider() { @Override @@ -231,6 +232,16 @@ public VectorFloat approximateCentroid() { }; } + private static Supplier asRandomAccessSupplier(Supplier>> supplier) { + return () -> { + var v = supplier.get(); + if (!(v instanceof RandomAccessVectorValues)) { + throw new IllegalStateException("Supplier returned VectorValues instance of " + v.getClass().getName() + " which does not implement RandomAccessVectorValues"); + } + return (RandomAccessVectorValues) v; + }; + } + /** * Returns a BSP that performs exact score comparisons using the given * {@link RandomAccessByteVectorValues} and {@link ByteVectorSimilarityFunction}. From f64773108903f2294e25bcdfac9fc74968d1dbfa Mon Sep 17 00:00:00 2001 From: Raghuveer Devulapalli Date: Fri, 25 Sep 2026 10:03:47 +0000 Subject: [PATCH 13/13] Support INT8 datasets in BenchYAML pipeline - Add SiftLoader.readBvecs() to parse .bvecs INT8 vector files - Parameterize DataSet to support FloatDataSet and ByteDataSet - Update DataSetLoaderSimpleMFD to load .bvecs files and resolve ByteVectorSimilarityFunction - Support INLINE_BYTE_VECTORS on-disk index format version 6 and in-memory graph builds - Support exact byte scoring in Grid.ConfiguredSystem and QueryExecutor - Adapt ThroughputBenchmark warmup queries for ByteDataSet - Validate that compression and reranking are disallowed for ByteDataSet in YAML configs - Add SIFT1M INT8 configuration to local-catalog.yaml, dataset-metadata.yml, and index parameters --- .../jvector/bench/CompactorBenchmark.java | 9 +- .../graph/disk/AbstractGraphIndexFormat.java | 2 +- .../jvector/example/AutoBenchYAML.java | 17 +- .../github/jbellis/jvector/example/Bench.java | 9 +- .../jbellis/jvector/example/BenchYAML.java | 8 +- .../jvector/example/CompactionBench.java | 10 +- .../github/jbellis/jvector/example/Grid.java | 248 +++++++++++++----- .../jvector/example/HelloVectorWorld.java | 6 +- .../example/benchmarks/QueryExecutor.java | 22 +- .../benchmarks/ThroughputBenchmark.java | 22 +- .../benchmarks/datasets/ByteDataSet.java | 107 ++++++++ .../example/benchmarks/datasets/DataSet.java | 55 ++-- .../benchmarks/datasets/DataSetInfo.java | 8 +- .../datasets/DataSetLoaderSimpleMFD.java | 21 ++ .../datasets/DataSetProperties.java | 13 + .../benchmarks/datasets/DataSetUtils.java | 26 +- .../{SimpleDataSet.java => FloatDataSet.java} | 49 ++-- .../example/reporting/DatasetInfoWriter.java | 4 +- .../example/reporting/RunArtifacts.java | 5 +- .../jvector/example/tutorial/DiskIntro.java | 4 +- .../example/tutorial/LargerThanMemory.java | 4 +- .../jvector/example/tutorial/NvqExample.java | 3 +- .../example/util/CompressorParameters.java | 18 +- .../example/util/DataSetPartitioner.java | 4 +- .../jvector/example/util/SiftLoader.java | 26 ++ .../example/yaml/CommonParameters.java | 22 +- .../jvector/example/yaml/Compression.java | 4 +- .../example/yaml/ConstructionParameters.java | 70 +++-- .../graph/disk/ParallelWriteExample.java | 4 +- .../datasets/DataSetPropertiesTest.java | 4 +- .../dataset-catalogs/local-catalog.yaml | 5 + .../yaml-configs/dataset-metadata.yml | 3 + .../sift1m-128-euclidean-int8.yml | 23 ++ .../jvector/microbench/GraphBuildBench.java | 6 +- 34 files changed, 601 insertions(+), 240 deletions(-) create mode 100644 jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/ByteDataSet.java rename jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/{SimpleDataSet.java => FloatDataSet.java} (79%) create mode 100644 jvector-examples/yaml-configs/index-parameters/sift1m-128-euclidean-int8.yml diff --git a/benchmarks-jmh/src/main/java/io/github/jbellis/jvector/bench/CompactorBenchmark.java b/benchmarks-jmh/src/main/java/io/github/jbellis/jvector/bench/CompactorBenchmark.java index 3a3b9c8bf..d66dc6b40 100644 --- a/benchmarks-jmh/src/main/java/io/github/jbellis/jvector/bench/CompactorBenchmark.java +++ b/benchmarks-jmh/src/main/java/io/github/jbellis/jvector/bench/CompactorBenchmark.java @@ -19,6 +19,7 @@ import io.github.jbellis.jvector.disk.ReaderSupplier; import io.github.jbellis.jvector.disk.ReaderSupplierFactory; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSetInfo; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.example.reporting.GitInfo; @@ -251,7 +252,7 @@ private static void writeCompletedCount(int count) { private List> queryVectors; private List> baseVectors; private List> groundTruth; - private DataSet ds; + private FloatDataSet ds; private VectorSimilarityFunction similarityFunction; private final List graphs = new ArrayList<>(); @@ -393,7 +394,7 @@ public void setup() throws Exception { boolean needsRecallData = measureRecall && workloadMode != WorkloadMode.PARTITION; if (needsBaseVectors) { - ds = DataSets.loadDataSet(datasetNames) + ds = (FloatDataSet) DataSets.loadDataSet(datasetNames) .orElseThrow(() -> new RuntimeException("Dataset not found: " + datasetNames)) .getDataSet(); @@ -432,7 +433,7 @@ public void setup() throws Exception { dimension = -1; if (needsRecallData) { - ds = DataSets.loadDataSet(datasetNames) + ds = (FloatDataSet) DataSets.loadDataSet(datasetNames) .orElseThrow(() -> new RuntimeException("Dataset not found: " + datasetNames)) .getDataSet(); queryVectors = ds.getQueryVectors(); @@ -595,7 +596,7 @@ private void verifyPartitionsExist(Path partitionsDir, int numPartitions) { } } - private void buildPartitions(DataSet ds, List> baseVectors) throws Exception { + private void buildPartitions(DataSet ds, List> baseVectors) throws Exception { var partitionedData = DataSetPartitioner.partition(baseVectors, numPartitions, splitDistribution); vectorsPerSourceCount = partitionedData.sizes; diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexFormat.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexFormat.java index d43497944..7d63a5a1e 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexFormat.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/disk/AbstractGraphIndexFormat.java @@ -414,7 +414,7 @@ public OnDiskGraphIndex loadOnDiskIndex(RandomAccessReader reader, Header header * because a new {@link FeatureId} is added to the enum for some future version. */ protected static Set allFeatures() { - return EnumSet.of(FeatureId.INLINE_VECTORS, FeatureId.FUSED_PQ, FeatureId.NVQ_VECTORS, + return EnumSet.of(FeatureId.INLINE_VECTORS, FeatureId.INLINE_BYTE_VECTORS, FeatureId.FUSED_PQ, FeatureId.NVQ_VECTORS, FeatureId.SEPARATED_VECTORS, FeatureId.SEPARATED_NVQ); } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/AutoBenchYAML.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/AutoBenchYAML.java index 24c39ae47..3986175a1 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/AutoBenchYAML.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/AutoBenchYAML.java @@ -21,6 +21,7 @@ import io.github.jbellis.jvector.example.util.BenchmarkSummarizer.SummaryStats; import io.github.jbellis.jvector.example.util.CheckpointManager; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.example.yaml.DatasetCollection; import io.github.jbellis.jvector.example.yaml.MultiConfig; @@ -123,7 +124,7 @@ public static void main(String[] args) throws IOException { logger.info("Loading dataset: {}", datasetName); try { - DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( + DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( () -> new RuntimeException("Dataset " + datasetName + " not found") ).getDataSet(); logger.info("Dataset loaded: {} with {} vectors", datasetName, ds.getBaseVectors().size()); @@ -148,15 +149,15 @@ public static void main(String[] args) throws IOException { List datasetResults = Grid.runAllAndCollectResults(ds, config.construction.useSavedIndexIfExists, - config.construction.outDegree, + config.construction.outDegree, config.construction.efConstruction, - config.construction.neighborOverflow, + config.construction.neighborOverflow, config.construction.addHierarchy, config.construction.refineFinalGraph, - config.construction.getFeatureSets(), - config.construction.getCompressorParameters(), - config.search.getCompressorParameters(), - config.search.topKOverquery, + config.construction.getFeatureSets(ds), + config.construction.getCompressorParameters(ds), + config.search.getCompressorParameters(ds), + config.search.topKOverquery, config.search.useSearchPruning); results.addAll(datasetResults); @@ -167,7 +168,7 @@ public static void main(String[] args) throws IOException { // Compaction regression — failures are non-fatal and don't block checkpointing try { logger.info("Running compaction benchmark for dataset: {}", datasetName); - List datasetCompactionResults = CompactionBench.run(ds); + List datasetCompactionResults = CompactionBench.run((FloatDataSet) ds); compactionResults.addAll(datasetCompactionResults); logger.info("Compaction benchmark completed for dataset: {} ({} configs)", datasetName, datasetCompactionResults.size()); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Bench.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Bench.java index 78a85e1fc..78e1fd8eb 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Bench.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Bench.java @@ -19,6 +19,7 @@ import io.github.jbellis.jvector.example.util.CompressorParameters; import io.github.jbellis.jvector.example.util.CompressorParameters.PQParameters; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.example.yaml.DatasetCollection; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; @@ -55,14 +56,14 @@ public static void main(String[] args) throws IOException { var addHierarchyGrid = List.of(true); // List.of(false, true); var refineFinalGraphGrid = List.of(true); // List.of(false, true); var usePruningGrid = List.of(true); // List.of(false, true); - List> buildCompression = Arrays.asList( + List> buildCompression = Arrays.asList( ds -> new PQParameters(ds.getDimension() / 8, 256, ds.getSimilarityFunction() == VectorSimilarityFunction.EUCLIDEAN, UNWEIGHTED), __ -> CompressorParameters.NONE ); - List> searchCompression = Arrays.asList( + List> searchCompression = Arrays.asList( __ -> CompressorParameters.NONE, // ds -> new CompressorParameters.BQParameters(), ds -> new PQParameters(ds.getDimension() / 8, @@ -85,13 +86,13 @@ public static void main(String[] args) throws IOException { execute(pattern, enableIndexCache, buildCompression, featureSets, searchCompression, mGrid, efConstructionGrid, neighborOverflowGrid, addHierarchyGrid, refineFinalGraphGrid, topKGrid, usePruningGrid); } - private static void execute(Pattern pattern, boolean enableIndexCache, List> buildCompression, List> featureSets, List> compressionGrid, List mGrid, List efConstructionGrid, List neighborOverflowGrid, List addHierarchyGrid, List refineFinalGraphGrid, Map> topKGrid, List usePruningGrid) throws IOException { + private static void execute(Pattern pattern, boolean enableIndexCache, List> buildCompression, List> featureSets, List> compressionGrid, List mGrid, List efConstructionGrid, List neighborOverflowGrid, List addHierarchyGrid, List refineFinalGraphGrid, Map> topKGrid, List usePruningGrid) throws IOException { var datasetCollection = DatasetCollection.load(); var datasetNames = datasetCollection.getAll().stream().filter(dn -> pattern.matcher(dn).find()).collect(Collectors.toList()); System.out.println("Executing the following datasets: " + datasetNames); for (var datasetName : datasetNames) { - DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( + FloatDataSet ds = (FloatDataSet) DataSets.loadDataSet(datasetName).orElseThrow( () -> new RuntimeException("Dataset " + datasetName + " not found") ).getDataSet(); Grid.runAll(ds, enableIndexCache, mGrid, efConstructionGrid, neighborOverflowGrid, addHierarchyGrid, refineFinalGraphGrid, featureSets, buildCompression, compressionGrid, topKGrid, usePruningGrid); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/BenchYAML.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/BenchYAML.java index e066a34dc..6263e3201 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/BenchYAML.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/BenchYAML.java @@ -118,7 +118,7 @@ public static void main(String[] args) throws IOException { for (var config : allConfigs) { String datasetName = config.dataset; - DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( + DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( () -> new RuntimeException("Could not load dataset:" + datasetName) ).getDataSet(); // Register dataset info the first time we actually load the dataset for benchmarking @@ -131,9 +131,9 @@ public static void main(String[] args) throws IOException { config.construction.neighborOverflow, config.construction.addHierarchy, config.construction.refineFinalGraph, - config.construction.getFeatureSets(), - config.construction.getCompressorParameters(), - config.search.getCompressorParameters(), + config.construction.getFeatureSets(ds), + config.construction.getCompressorParameters(ds), + config.search.getCompressorParameters(ds), config.search.topKOverquery, config.search.useSearchPruning, artifacts); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/CompactionBench.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/CompactionBench.java index 15543ebbc..5cfd9672d 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/CompactionBench.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/CompactionBench.java @@ -18,7 +18,7 @@ import io.github.jbellis.jvector.disk.ReaderSupplier; import io.github.jbellis.jvector.disk.ReaderSupplierFactory; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.util.AccuracyMetrics; import io.github.jbellis.jvector.example.util.CompactionPartitionSource; import io.github.jbellis.jvector.example.yaml.TestDataPartition.Distribution; @@ -97,7 +97,7 @@ private CompactionBench() {} * one result per config. A config that fails (e.g. missing partitions) is logged and skipped so * the remaining configs still run. Throws if the dataset has no query vectors or ground truth. */ - public static List run(DataSet ds) throws Exception { + public static List run(FloatDataSet ds) throws Exception { var queryVectors = ds.getQueryVectors(); var groundTruth = ds.getGroundTruth(); if (queryVectors == null || queryVectors.isEmpty()) { @@ -118,7 +118,7 @@ public static List run(DataSet ds) throws Exception { return results; } - private static BenchResult runConfig(DataSet ds, PartitionConfig cfg) throws Exception { + private static BenchResult runConfig(FloatDataSet ds, PartitionConfig cfg) throws Exception { String datasetName = ds.getName(); logger.info("Compaction bench [{}] config {}: {} vectors", datasetName, cfg.dirName(), ds.getBaseVectors().size()); @@ -135,7 +135,7 @@ private static BenchResult runConfig(DataSet ds, PartitionConfig cfg) throws Exc } } - private static BenchResult compactAndMeasure(DataSet ds, PartitionConfig cfg, + private static BenchResult compactAndMeasure(FloatDataSet ds, PartitionConfig cfg, List partitionPaths, Path tempDir) throws Exception { List> baseVectors = ds.getBaseVectors(); int dimension = ds.getDimension(); @@ -243,7 +243,7 @@ static final class SearchStats { * Searches every query against the compacted graph, timing each search, and returns recall plus * mean and p99 per-query latency (ms) and throughput (queries/sec, single-threaded sequential). */ - private static SearchStats searchCompacted(Path indexPath, DataSet ds, + private static SearchStats searchCompacted(Path indexPath, FloatDataSet ds, List> baseVectors, int dimension, VectorSimilarityFunction vsf) throws Exception { var queryVectors = ds.getQueryVectors(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Grid.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Grid.java index 548fd4b47..11e84f063 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Grid.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/Grid.java @@ -28,7 +28,9 @@ import io.github.jbellis.jvector.example.benchmarks.ThroughputBenchmark; import io.github.jbellis.jvector.example.benchmarks.diagnostics.BenchmarkDiagnostics; import io.github.jbellis.jvector.example.benchmarks.diagnostics.DiagnosticLevel; +import io.github.jbellis.jvector.example.benchmarks.datasets.ByteDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.reporting.*; import io.github.jbellis.jvector.example.reporting.RunArtifacts; import io.github.jbellis.jvector.example.util.CompressorParameters; @@ -38,11 +40,14 @@ import io.github.jbellis.jvector.graph.ImmutableGraphIndex; import io.github.jbellis.jvector.graph.GraphIndexBuilder; import io.github.jbellis.jvector.graph.GraphSearcher; +import io.github.jbellis.jvector.graph.ListRandomAccessByteVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessByteVectorValues; import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import io.github.jbellis.jvector.graph.disk.*; import io.github.jbellis.jvector.graph.disk.feature.Feature; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; import io.github.jbellis.jvector.graph.disk.feature.FusedPQ; +import io.github.jbellis.jvector.graph.disk.feature.InlineByteVectors; import io.github.jbellis.jvector.graph.disk.feature.InlineVectors; import io.github.jbellis.jvector.graph.disk.feature.NVQ; import io.github.jbellis.jvector.graph.similarity.BuildScoreProvider; @@ -56,6 +61,8 @@ import io.github.jbellis.jvector.quantization.VectorCompressor; import io.github.jbellis.jvector.util.ExplicitThreadLocal; import io.github.jbellis.jvector.util.PhysicalCoreExecutor; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import java.io.FileNotFoundException; @@ -94,7 +101,7 @@ public static Double getIndexBuildTimeSeconds(String datasetName) { return indexBuildTimes.get(datasetName); } - static void runAll(DataSet ds, + static void runAll(DataSet ds, boolean enableIndexCache, List mGrid, List efConstructionGrid, @@ -102,8 +109,8 @@ static void runAll(DataSet ds, List addHierarchyGrid, List refineFinalGraphGrid, List> featureSets, - List> buildCompressors, - List> compressionGrid, + List> buildCompressors, + List> compressionGrid, Map> topKGrid, List usePruningGrid, RunArtifacts artifacts @@ -158,7 +165,7 @@ static void runAll(DataSet ds, } // Overload for legacy callers that do not use yaml-configs - static void runAll(DataSet ds, + static void runAll(DataSet ds, boolean enableIndexCache, List mGrid, List efConstructionGrid, @@ -166,8 +173,8 @@ static void runAll(DataSet ds, List addHierarchyGrid, List refineFinalGraphGrid, List> featureSets, - List> buildCompressors, - List> compressionGrid, + List> buildCompressors, + List> compressionGrid, Map> topKGrid, List usePruningGrid) throws IOException { @@ -218,12 +225,12 @@ static void runOneGraph(OnDiskGraphIndexCache cache, float neighborOverflow, boolean addHierarchy, boolean refineFinalGraph, - Function buildCompressor, - List> compressionGrid, + Function buildCompressor, + List> compressionGrid, Map> topKGrid, List usePruningGrid, RunArtifacts artifacts, - DataSet ds, + DataSet ds, Path workDirectory) throws IOException { // Prepare to collect index construction metrics for reporting.... @@ -234,16 +241,22 @@ static void runOneGraph(OnDiskGraphIndexCache cache, diagnostics.startMonitoring("testDirectory", workDirectory); diagnostics.startMonitoring("indexCache", Paths.get(indexCacheDir)); diagnostics.capturePrePhaseSnapshot("Graph Build"); - System.out.printf("%s: Dataset similarity function is %s%n", ds.getName(), ds.getSimilarityFunction()); + FloatDataSet floatDs = ds instanceof FloatDataSet ? (FloatDataSet) ds : null; + ByteDataSet byteDs = ds instanceof ByteDataSet ? (ByteDataSet) ds : null; + if (floatDs != null) { + System.out.printf("%s: Dataset similarity function is %s%n", ds.getName(), floatDs.getSimilarityFunction()); + } else if (byteDs != null) { + System.out.printf("%s: Dataset similarity function is %s%n", ds.getName(), byteDs.getByteSimilarityFunction()); + } // Resolve build compressor (and label quant type) so we can record compute time VectorCompressor buildCompressorObj = null; String buildQuantType = null; - if (buildCompressor != null) { - var buildParams = buildCompressor.apply(ds); + if (buildCompressor != null && floatDs != null) { + var buildParams = buildCompressor.apply(floatDs); buildQuantType = quantTypeOf(buildParams); // "PQ", "BQ", or null - buildCompressorObj = getCompressor(buildCompressor, ds, constructionMetrics, Phase.INDEX, buildQuantType); + buildCompressorObj = getCompressor(buildCompressor, floatDs, constructionMetrics, Phase.INDEX, buildQuantType); } // Used in logging @@ -288,7 +301,7 @@ static void runOneGraph(OnDiskGraphIndexCache cache, // At least one index needs to be built (b/c not in cache or cache is disabled) // We pass the handles map so buildOnDisk knows exactly where to write var result = buildOnDisk(missing, M, efConstruction, neighborOverflow, addHierarchy, refineFinalGraph, - ds, outputDir, buildCompressorObj, handles, constructionMetrics); + floatDs, outputDir, buildCompressorObj, handles, constructionMetrics); indexes.putAll(result.indexes); indexFileSizes.putAll(result.fileSizes); } @@ -314,24 +327,28 @@ static void runOneGraph(OnDiskGraphIndexCache cache, } else { constructionMetrics.resetSearch(); // per (index, cpSupplier) config - var searchParams = cpSupplier.apply(ds); - String searchQuantType = quantTypeOf(searchParams); // "PQ", "BQ", or null - var compressor = getCompressor(cpSupplier, ds, constructionMetrics, Phase.SEARCH, searchQuantType); - - if (compressor == null) { + if (floatDs == null) { cv = null; - System.out.format("%s: No search compressor configured, FULL PRECISION vectors will be used for search%n", ds.getName()); } else { - long start = System.nanoTime(); - cv = constructionMetrics.search(searchQuantType) - .timeEncode(() -> compressor.encodeAll(ds.getBaseRavv())); - double encodingTimeS = (System.nanoTime() - start) / 1_000_000_000.0; - if (cv == null) { - throw new IllegalStateException(String.format( - "Compressor '%s' was provided but failed to encode vectors for dataset '%s'. " + - "Aborting to prevent false recall results.", compressor, ds.getName())); + var searchParams = cpSupplier.apply(floatDs); + String searchQuantType = quantTypeOf(searchParams); // "PQ", "BQ", or null + var compressor = getCompressor(cpSupplier, floatDs, constructionMetrics, Phase.SEARCH, searchQuantType); + + if (compressor == null) { + cv = null; + System.out.format("%s: No search compressor configured, FULL PRECISION vectors will be used for search%n", ds.getName()); + } else { + long start = System.nanoTime(); + cv = constructionMetrics.search(searchQuantType) + .timeEncode(() -> compressor.encodeAll(floatDs.getBaseRavv())); + double encodingTimeS = (System.nanoTime() - start) / 1_000_000_000.0; + if (cv == null) { + throw new IllegalStateException(String.format( + "Compressor '%s' was provided but failed to encode vectors for dataset '%s'. " + + "Aborting to prevent false recall results.", compressor, ds.getName())); + } + System.out.format("%s: %s encoded %d vectors [%.2f MB] in %.2fs%n", ds.getName(), compressor, floatDs.getBaseVectors().size(), (cv.ramBytesUsed() / 1024f / 1024f), encodingTimeS); } - System.out.format("%s: %s encoded %d vectors [%.2f MB] in %.2fs%n", ds.getName(), compressor, ds.getBaseVectors().size(), (cv.ramBytesUsed() / 1024f / 1024f), encodingTimeS); } } @@ -385,7 +402,7 @@ private static BuildOnDiskResult buildOnDisk(List> feat float neighborOverflow, boolean addHierarchy, boolean refineFinalGraph, - DataSet ds, + FloatDataSet ds, Path outputDir, VectorCompressor buildCompressor, Map, OnDiskGraphIndexCache.WriteHandle> handles, @@ -503,7 +520,21 @@ private static BuilderWithSuppliers builderWithSuppliers(Set features ConstructionMetrics constructionMetrics) throws FileNotFoundException { - var identityMapper = new OrdinalMapper.IdentityMapper(floatVectors.size() - 1); + return builderWithSuppliers(features, onHeapGraph, outPath, floatVectors, null, pq, constructionMetrics); + } + + private static BuilderWithSuppliers builderWithSuppliers(Set features, + ImmutableGraphIndex onHeapGraph, + Path outPath, + RandomAccessVectorValues floatVectors, + RandomAccessByteVectorValues byteVectors, + ProductQuantization pq, + ConstructionMetrics constructionMetrics) + throws FileNotFoundException + { + int dimension = floatVectors != null ? floatVectors.dimension() : byteVectors.dimension(); + int size = floatVectors != null ? floatVectors.size() : byteVectors.size(); + var identityMapper = new OrdinalMapper.IdentityMapper(size - 1); var builder = new RandomAccessOnDiskGraphIndexWriter.Builder(onHeapGraph, outPath); builder.withMapper(identityMapper); @@ -511,9 +542,16 @@ private static BuilderWithSuppliers builderWithSuppliers(Set features for (var featureId : features) { switch (featureId) { case INLINE_VECTORS: - builder.with(new InlineVectors(floatVectors.dimension())); + builder.with(new InlineVectors(dimension)); suppliers.put(FeatureId.INLINE_VECTORS, ordinal -> new InlineVectors.State(floatVectors.getVector(ordinal))); break; + case INLINE_BYTE_VECTORS: + if (byteVectors == null) { + throw new IllegalArgumentException("INLINE_BYTE_VECTORS requested but no byte vectors provided"); + } + builder.with(new InlineByteVectors(dimension)); + suppliers.put(FeatureId.INLINE_BYTE_VECTORS, ordinal -> new InlineByteVectors.State(byteVectors.getVector(ordinal))); + break; case FUSED_PQ: if (pq == null) { System.out.println("Skipping Fused ADC feature due to null ProductQuantization"); @@ -523,14 +561,16 @@ private static BuilderWithSuppliers builderWithSuppliers(Set features builder.with(new FusedPQ(onHeapGraph.maxDegree(), pq)); break; case NVQ_VECTORS: - int nSubVectors = floatVectors.dimension() == 2 ? 1 : 2; + if (byteVectors != null) { + throw new IllegalArgumentException("NVQ_VECTORS is not supported for INT8 (byte) datasets"); + } + int nSubVectors = dimension == 2 ? 1 : 2; var nvq = (constructionMetrics != null) ? constructionMetrics.index("NVQ").timeCompute(() -> NVQuantization.compute(floatVectors, nSubVectors)) : NVQuantization.compute(floatVectors, nSubVectors); builder.with(new NVQ(nvq)); suppliers.put(FeatureId.NVQ_VECTORS, ordinal -> new NVQ.State(nvq.encode(floatVectors.getVector(ordinal)))); break; - } } return new BuilderWithSuppliers(builder, suppliers); @@ -571,9 +611,27 @@ private static Map, ImmutableGraphIndex> buildInMemory(List ds, Path testDirectory) throws IOException + { + if (ds instanceof ByteDataSet) { + return buildInMemoryByte(featureSets, M, efConstruction, neighborOverflow, addHierarchy, refineFinalGraph, + (ByteDataSet) ds, testDirectory); + } + return buildInMemoryFloat(featureSets, M, efConstruction, neighborOverflow, addHierarchy, refineFinalGraph, + (FloatDataSet) ds, testDirectory); + } + + private static Map, ImmutableGraphIndex> buildInMemoryFloat(List> featureSets, + int M, + int efConstruction, + float neighborOverflow, + boolean addHierarchy, + boolean refineFinalGraph, + FloatDataSet ds, + Path testDirectory) + throws IOException { var floatVectors = ds.getBaseRavv(); Map, ImmutableGraphIndex> indexes = new HashMap<>(); @@ -593,16 +651,10 @@ private static Map, ImmutableGraphIndex> buildInMemory(List, ImmutableGraphIndex> buildInMemory(List, ImmutableGraphIndex> buildInMemory(List, ImmutableGraphIndex> buildInMemoryByte(List> featureSets, + int M, + int efConstruction, + float neighborOverflow, + boolean addHierarchy, + boolean refineFinalGraph, + ByteDataSet ds, + Path testDirectory) + throws IOException + { + for (var features : featureSets) { + if (features.contains(FeatureId.FUSED_PQ) || features.contains(FeatureId.NVQ_VECTORS)) { + throw new IllegalArgumentException( + "FUSED_PQ and NVQ_VECTORS are not supported for INT8 datasets; got: " + features); + } + } + + var rabvv = ds.getBaseByteRavv(); + ByteVectorSimilarityFunction bvsf = ds.getByteSimilarityFunction(); + + Map, ImmutableGraphIndex> indexes = new HashMap<>(); + long start; + + try (GraphIndexBuilder builder = GraphIndexBuilder.builder(rabvv, bvsf, M) + .withBeamWidth(efConstruction) + .withNeighborOverflow(neighborOverflow) + .withAlpha(1.2f) + .withSimdExecutor(PhysicalCoreExecutor.pool()) + .withParallelExecutor(FilteredForkJoinPool.createFilteredPool()) + .build()) + { + start = System.nanoTime(); + var onHeapGraph = builder.build(rabvv); + double buildTimeS = (System.nanoTime() - start) / 1_000_000_000.0; + System.out.format("Build (INT8) M=%d overflow=%.2f ef=%d in %.2fs%n", + M, neighborOverflow, efConstruction, buildTimeS); + for (int i = 0; i <= onHeapGraph.getMaxLevel(); i++) { + System.out.format(" L%d: %d nodes, %.2f avg degree%n", + i, onHeapGraph.size(i), onHeapGraph.getAverageDegree(i)); + } + + int n = 0; + for (var features : featureSets) { + var graphPath = testDirectory.resolve("graph" + n++); + var bws = builderWithSuppliers(features, onHeapGraph, graphPath, null, rabvv, null, null); + try (var writer = bws.builder.build()) { + start = System.nanoTime(); + writer.write(bws.suppliers); + System.out.format("Wrote %s in %.2fs%n", features, (System.nanoTime() - start) / 1_000_000_000.0); + } + var index = OnDiskGraphIndex.load(ReaderSupplierFactory.open(graphPath)); + indexes.put(features, index); + } + indexBuildTimes.put(ds.getName(), buildTimeS); + } + return indexes; + } + // avoid recomputing the compressor repeatedly (this is a relatively small memory footprint) static final Map> cachedCompressors = new IdentityHashMap<>(); @@ -827,7 +936,7 @@ private static Map ordered(Object... kv) { } public static List runAllAndCollectResults( - DataSet ds, + DataSet ds, boolean enableIndexCache, List mGrid, List efConstructionGrid, @@ -835,8 +944,8 @@ public static List runAllAndCollectResults( List addHierarchyGrid, List refineFinalGraphGrid, List> featureSets, - List> buildCompressors, - List> compressionGrid, + List> buildCompressors, + List> compressionGrid, Map> topKGrid, List usePruningGrid) throws IOException { @@ -853,8 +962,8 @@ public static List runAllAndCollectResults( for (boolean addHierarchy : addHierarchyGrid) { for (boolean refineFinalGraph : refineFinalGraphGrid) { for (Set features : featureSets) { - for (Function buildCompressor : buildCompressors) { - for (Function searchCompressor : compressionGrid) { + for (Function buildCompressor : buildCompressors) { + for (Function searchCompressor : compressionGrid) { Path testDirectory = Files.createTempDirectory("bench"); try (var diagnostics = new BenchmarkDiagnostics(getDiagnosticLevel())) { // Capture initial state @@ -864,8 +973,9 @@ public static List runAllAndCollectResults( Map, ImmutableGraphIndex> indexes = new HashMap<>(); Map, Long> indexFileSizes = new HashMap<>(); - var compressor = getCompressor(buildCompressor, ds); - var searchCompressorObj = getCompressor(searchCompressor, ds); + var floatDs = (FloatDataSet) ds; + var compressor = getCompressor(buildCompressor, floatDs); + var searchCompressorObj = getCompressor(searchCompressor, floatDs); // Encode vectors for reranking if a compressor is provided CompressedVectors cvArg; if (features.contains(FeatureId.FUSED_PQ)) { @@ -877,7 +987,7 @@ public static List runAllAndCollectResults( System.out.format("%s: No search compressor configured, " + "FULL PRECISION vectors will be used for search%n", ds.getName()); } else { - cvArg = searchCompressorObj.encodeAll(ds.getBaseRavv()); + cvArg = searchCompressorObj.encodeAll(floatDs.getBaseRavv()); if (cvArg == null) { throw new IllegalStateException(String.format( "Compressor '%s' was provided but failed to encode vectors for dataset '%s'. " + @@ -885,7 +995,7 @@ public static List runAllAndCollectResults( searchCompressorObj, ds.getName())); } System.out.format("%s: %s encoded %d vectors [%.2f MB] for search%n", - ds.getName(), searchCompressorObj, ds.getBaseVectors().size(), + ds.getName(), searchCompressorObj, floatDs.getBaseVectors().size(), (cvArg.ramBytesUsed() / 1024f / 1024f)); } } @@ -923,7 +1033,7 @@ public static List runAllAndCollectResults( // At least one index needs to be built (b/c not in cache or cache is disabled) // We pass the handles map so buildOnDisk knows exactly where to write var result = buildOnDisk(missing, m, ef, neighborOverflow, addHierarchy, refineFinalGraph, - ds, outputDir, compressor, handles, null); + floatDs, outputDir, compressor, handles, null); indexes.putAll(result.indexes); indexFileSizes.putAll(result.fileSizes); } @@ -1024,7 +1134,7 @@ public static List runAllAndCollectResults( } /** Overload for non-reporting use */ - private static VectorCompressor getCompressor(Function cpSupplier, DataSet ds) { + private static VectorCompressor getCompressor(Function cpSupplier, FloatDataSet ds) { return getCompressor(cpSupplier, ds, null, null, null); } @@ -1045,8 +1155,8 @@ private static VectorCompressor getCompressor(Function getCompressor(Function cpSupplier, - DataSet ds, + private static VectorCompressor getCompressor(Function cpSupplier, + FloatDataSet ds, ConstructionMetrics metrics, Phase phase, String quantType) { @@ -1091,7 +1201,7 @@ private static VectorCompressor getCompressor(Function loadFromCache(DataSet ds, String fname, Path path) { + private static VectorCompressor loadFromCache(DataSet ds, String fname, Path path) { try (var readerSupplier = ReaderSupplierFactory.open(path); var rar = readerSupplier.get()) { var pq = ProductQuantization.load(rar); @@ -1115,7 +1225,7 @@ private static void saveToCache(String fname, VectorCompressor compressor) { } public static class ConfiguredSystem implements AutoCloseable { - DataSet ds; + DataSet ds; ImmutableGraphIndex index; CompressedVectors cv; Set features; @@ -1124,7 +1234,7 @@ public static class ConfiguredSystem implements AutoCloseable { return new GraphSearcher(index); }); - ConfiguredSystem(DataSet ds, ImmutableGraphIndex index, CompressedVectors cv, Set features) { + ConfiguredSystem(DataSet ds, ImmutableGraphIndex index, CompressedVectors cv, Set features) { this.ds = ds; this.index = index; this.cv = cv; @@ -1134,25 +1244,33 @@ public static class ConfiguredSystem implements AutoCloseable { public SearchScoreProvider scoreProviderFor(VectorFloat queryVector, ImmutableGraphIndex.View view) { var scoringView = (ImmutableGraphIndex.ScoringView) view; ScoreFunction.ApproximateScoreFunction asf; + var floatDs = (FloatDataSet) ds; if (features.contains(FeatureId.FUSED_PQ)) { - asf = scoringView.approximateScoreFunctionFor(queryVector, ds.getSimilarityFunction()); + asf = scoringView.approximateScoreFunctionFor(queryVector, floatDs.getSimilarityFunction()); } else { // if we're not compressing then just use the exact score function if (cv == null) { - return DefaultSearchScoreProvider.exact(queryVector, ds.getSimilarityFunction(), ds.getBaseRavv()); + return DefaultSearchScoreProvider.exact(queryVector, floatDs.getSimilarityFunction(), floatDs.getBaseRavv()); } - asf = cv.precomputedScoreFunctionFor(queryVector, ds.getSimilarityFunction()); + asf = cv.precomputedScoreFunctionFor(queryVector, floatDs.getSimilarityFunction()); } - var rr = scoringView.rerankerFor(queryVector, ds.getSimilarityFunction()); + var rr = scoringView.rerankerFor(queryVector, floatDs.getSimilarityFunction()); return new DefaultSearchScoreProvider(asf, rr); } + public SearchScoreProvider scoreProviderFor(ByteSequence queryBytes, ImmutableGraphIndex.View view) { + var diskView = (OnDiskGraphIndex.View) view; + var byteDs = (ByteDataSet) ds; + var reranker = diskView.byteVectorRerankerFor(queryBytes, byteDs.getByteSimilarityFunction()); + return new DefaultSearchScoreProvider(reranker); + } + public GraphSearcher getSearcher() { return searchers.get(); } - public DataSet getDataSet() { + public DataSet getDataSet() { return ds; } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/HelloVectorWorld.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/HelloVectorWorld.java index ea4752e4b..b5eee0679 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/HelloVectorWorld.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/HelloVectorWorld.java @@ -52,9 +52,9 @@ public static void main(String[] args) throws IOException { config.construction.neighborOverflow, config.construction.addHierarchy, config.construction.refineFinalGraph, - config.construction.getFeatureSets(), - config.construction.getCompressorParameters(), - config.search.getCompressorParameters(), + config.construction.getFeatureSets(ds), + config.construction.getCompressorParameters(ds), + config.search.getCompressorParameters(ds), config.search.topKOverquery, config.search.useSearchPruning, artifacts); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/QueryExecutor.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/QueryExecutor.java index 45094bbd6..2e04b3f98 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/QueryExecutor.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/QueryExecutor.java @@ -17,8 +17,11 @@ package io.github.jbellis.jvector.example.benchmarks; import io.github.jbellis.jvector.example.Grid.ConfiguredSystem; +import io.github.jbellis.jvector.example.benchmarks.datasets.ByteDataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.graph.SearchResult; import io.github.jbellis.jvector.util.Bits; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; public class QueryExecutor { @@ -33,13 +36,26 @@ public class QueryExecutor { * @return the SearchResult for query i. */ public static SearchResult executeQuery(ConfiguredSystem cs, int topK, int rerankK, boolean usePruning, int i) { - var queryVector = cs.getDataSet().getQueryVectors().get(i); + if (cs.getDataSet() instanceof ByteDataSet) { + var queryBytes = ((ByteDataSet) cs.getDataSet()).getQueryVectors().get(i); + var searcher = cs.getSearcher(); + searcher.usePruning(usePruning); + var sf = cs.scoreProviderFor(queryBytes, searcher.getView()); + return searcher.search(sf, topK, rerankK, 0.0f, 0.0f, Bits.ALL); + } + var queryVector = ((FloatDataSet) cs.getDataSet()).getQueryVectors().get(i); return executeQuery(cs, topK, rerankK, usePruning, queryVector); } // Overload to allow single query injection (e.g., for warm-up with random vectors) - public static SearchResult executeQuery(ConfiguredSystem cs, int topK, int rerankK, boolean usePruning, VectorFloat queryVector - ) { + public static SearchResult executeQuery(ConfiguredSystem cs, int topK, int rerankK, boolean usePruning, VectorFloat queryVector) { + var searcher = cs.getSearcher(); + searcher.usePruning(usePruning); + var sf = cs.scoreProviderFor(queryVector, searcher.getView()); + return searcher.search(sf, topK, rerankK, 0.0f, 0.0f, Bits.ALL); + } + + public static SearchResult executeQuery(ConfiguredSystem cs, int topK, int rerankK, boolean usePruning, ByteSequence queryVector) { var searcher = cs.getSearcher(); searcher.usePruning(usePruning); var sf = cs.scoreProviderFor(queryVector, searcher.getView()); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/ThroughputBenchmark.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/ThroughputBenchmark.java index eaca57a17..c8823baa9 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/ThroughputBenchmark.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/ThroughputBenchmark.java @@ -154,14 +154,22 @@ public List runBenchmark( .forEach(k -> { long queryStart = System.nanoTime(); - // Generate a random vector - VectorFloat randQ = vts.createFloatVector(dim); - for (int j = 0; j < dim; j++) { - randQ.set(j, ThreadLocalRandom.current().nextFloat()); + SearchResult sr; + if (cs.getDataSet() instanceof io.github.jbellis.jvector.example.benchmarks.datasets.ByteDataSet) { + var randQ = vts.createByteSequence(dim); + for (int j = 0; j < dim; j++) { + randQ.set(j, (byte) ThreadLocalRandom.current().nextInt(-128, 128)); + } + sr = QueryExecutor.executeQuery(cs, topK, rerankK, usePruning, randQ); + } else { + // Generate a random vector + VectorFloat randQ = vts.createFloatVector(dim); + for (int j = 0; j < dim; j++) { + randQ.set(j, ThreadLocalRandom.current().nextFloat()); + } + VectorUtil.l2normalize(randQ); + sr = QueryExecutor.executeQuery(cs, topK, rerankK, usePruning, randQ); } - VectorUtil.l2normalize(randQ); - SearchResult sr = QueryExecutor.executeQuery( - cs, topK, rerankK, usePruning, randQ); SINK += sr.getVisitedCount(); long queryEnd = System.nanoTime(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/ByteDataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/ByteDataSet.java new file mode 100644 index 000000000..caed03e37 --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/ByteDataSet.java @@ -0,0 +1,107 @@ +/* + * Copyright DataStax, Inc. + * + * Licensed 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 io.github.jbellis.jvector.example.benchmarks.datasets; + +import io.github.jbellis.jvector.graph.ListRandomAccessByteVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessByteVectorValues; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; +import io.github.jbellis.jvector.vector.types.ByteSequence; + +import java.util.List; + +/** + * An int8 dataset backed by in-memory {@link ByteSequence} lists. + * + *

Type-specific accessors ({@link #getBaseByteRavv()}, {@link #getByteSimilarityFunction()}) + * live here rather than on the {@link DataSet} interface. + */ +public class ByteDataSet implements DataSet> { + private final String name; + private final ByteVectorSimilarityFunction similarityFunction; + private final List> baseVectors; + private final List> queryVectors; + private final List> groundTruth; + private RandomAccessByteVectorValues baseRavv; + + public ByteDataSet(String name, + ByteVectorSimilarityFunction similarityFunction, + List> baseVectors, + List> queryVectors, + List> groundTruth) + { + if (baseVectors.isEmpty()) { + throw new IllegalArgumentException("Base vectors must not be empty"); + } + if (queryVectors.isEmpty()) { + throw new IllegalArgumentException("Query vectors must not be empty"); + } + if (groundTruth.isEmpty()) { + throw new IllegalArgumentException("Ground truth must not be empty"); + } + if (baseVectors.get(0).length() != queryVectors.get(0).length()) { + throw new IllegalArgumentException("Base and query vectors must have the same dimensionality"); + } + if (queryVectors.size() != groundTruth.size()) { + throw new IllegalArgumentException("Query and ground truth lists must be the same size"); + } + + this.name = name; + this.similarityFunction = similarityFunction; + this.baseVectors = baseVectors; + this.queryVectors = queryVectors; + this.groundTruth = groundTruth; + + System.out.format("%n%s: %d base and %d query vectors created, dimensions %d%n", + name, baseVectors.size(), queryVectors.size(), baseVectors.get(0).length()); + } + + @Override + public String getName() { + return name; + } + + @Override + public int getDimension() { + return baseVectors.get(0).length(); + } + + @Override + public List> getBaseVectors() { + return baseVectors; + } + + @Override + public List> getQueryVectors() { + return queryVectors; + } + + @Override + public List> getGroundTruth() { + return groundTruth; + } + + public RandomAccessByteVectorValues getBaseByteRavv() { + if (baseRavv == null) { + baseRavv = new ListRandomAccessByteVectorValues(baseVectors, getDimension()); + } + return baseRavv; + } + + public ByteVectorSimilarityFunction getByteSimilarityFunction() { + return similarityFunction; + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSet.java index a33a40d31..63c3147d3 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSet.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSet.java @@ -16,61 +16,44 @@ package io.github.jbellis.jvector.example.benchmarks.datasets; -import io.github.jbellis.jvector.graph.RandomAccessVectorValues; -import io.github.jbellis.jvector.vector.VectorSimilarityFunction; -import io.github.jbellis.jvector.vector.types.VectorFloat; - -import java.util.*; +import java.util.List; /** - * This provides a uniform way to access vector test data, regardless of where it comes from or how it is implemented. + * Uniform access to vector test data, regardless of element type ({@code VectorFloat} for + * float32 datasets, {@code ByteSequence} for int8 datasets). + * + *

Type-specific accessors (RAVV, similarity function) live on the concrete subclasses + * {@link FloatDataSet} and {@link ByteDataSet} rather than here. + * + * @param the vector element type */ -public interface DataSet { - - /** - * Get dimensions of the vectors in this dataset. - * @return the dimensionality - */ - int getDimension(); - - /** - * Get a random-access view of base vectors. - * @return base vectors - */ - RandomAccessVectorValues getBaseRavv(); +public interface DataSet { /** * The symbolic name of this dataset, used for dataset selection and result labeling. - * @return the dataset name */ String getName(); /** - * The similarity function originally used to build this dataset, and the one that should be used for testing - * during indexing and traversal. - * @return the similarity function + * Dimensionality of the vectors in this dataset. */ - VectorSimilarityFunction getSimilarityFunction(); + int getDimension(); /** - * The base vectors as a list. - * @return a list of base vectors + * Base vectors as a list. */ - List> getBaseVectors(); + List getBaseVectors(); /** - * The query vectors as a list. - * Each major index corresponds to the self-same index from {@link #getGroundTruth()}. - * Ideally, the query vectors are disjoint with respect to the base vectors to improve testing integrity. - * @return a list of query vectors + * Query vectors as a list. + * Each index corresponds to the same index in {@link #getGroundTruth()}. */ - List> getQueryVectors(); + List getQueryVectors(); /** - * The ground truth as a list. - * Each major index corresponds to the self-same index from {@link #getQueryVectors()}. - * Each minor index within represents the corresponding ordinal from {@link #getBaseVectors()} and {@link #getBaseRavv()}. - * @return a list of query vectors. + * Ground truth as a list of neighbor-ordinal lists. + * Each major index corresponds to the same index in {@link #getQueryVectors()}. + * Each minor index is an ordinal into {@link #getBaseVectors()}. */ List> getGroundTruth(); } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetInfo.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetInfo.java index 94eb9f011..60af5a59b 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetInfo.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetInfo.java @@ -49,9 +49,9 @@ /// @see DataSetLoader /// @see DataSets public class DataSetInfo implements DataSetProperties { - private final Supplier loader; + private final Supplier> loader; private final DataSetProperties baseProperties; - private volatile DataSet cached; + private volatile DataSet cached; /// Creates a new dataset info handle. /// @@ -61,7 +61,7 @@ public class DataSetInfo implements DataSetProperties { /// /// @param baseProperties the dataset properties (name, similarity function, etc.) /// @param loader a supplier that performs the deferred load; invoked at most once - public DataSetInfo(DataSetProperties baseProperties, Supplier loader) { + public DataSetInfo(DataSetProperties baseProperties, Supplier> loader) { this.baseProperties = baseProperties; this.loader = loader; } @@ -124,7 +124,7 @@ public boolean isDuplicateVectorFree() { /// completes, after which all callers share the same cached instance. /// /// @return the ready-to-use {@link DataSet} - public DataSet getDataSet() { + public DataSet getDataSet() { if (cached == null) { synchronized (this) { if (cached == null) { diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFD.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFD.java index 5582e27e8..12eefd0b5 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFD.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFD.java @@ -15,6 +15,7 @@ */ package io.github.jbellis.jvector.example.benchmarks.datasets; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.example.util.SiftLoader; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -423,6 +424,26 @@ public Optional loadDataSet(String dataSetName) { String.format( "Dataset '%s' was found in dataset catalog, but no metadata entry was found in dataset-metadata.yml. ", dataSetName))); + boolean baseIsBvecs = baseFile.endsWith(".bvecs"); + boolean queryIsBvecs = queryFile.endsWith(".bvecs"); + if (baseIsBvecs != queryIsBvecs) { + throw new IllegalArgumentException( + "Dataset '" + dataSetName + "': base and query files must use the same format, " + + "but got '" + baseFile + "' and '" + queryFile + "'"); + } + + if (baseIsBvecs) { + ByteVectorSimilarityFunction bvsf = props.byteSimilarityFunction() + .orElseThrow(() -> new IllegalArgumentException( + "Dataset '" + dataSetName + "' uses .bvecs files but has no similarity_function configured")); + return Optional.of(new DataSetInfo(props, () -> { + var baseVectors = SiftLoader.readBvecs(effectiveCacheDir.resolve(baseFile).toString()); + var queryVectors = SiftLoader.readBvecs(effectiveCacheDir.resolve(queryFile).toString()); + var gtVectors = SiftLoader.readIvecs(effectiveCacheDir.resolve(gtFile).toString()); + return new ByteDataSet(dataSetName, bvsf, baseVectors, queryVectors, gtVectors); + })); + } + return Optional.of(new DataSetInfo(props, () -> { var baseVectors = SiftLoader.readFvecs(effectiveCacheDir.resolve(baseFile).toString()); var queryVectors = SiftLoader.readFvecs(effectiveCacheDir.resolve(queryFile).toString()); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetProperties.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetProperties.java index 6d017d899..8df80c076 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetProperties.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetProperties.java @@ -16,6 +16,7 @@ package io.github.jbellis.jvector.example.benchmarks.datasets; +import io.github.jbellis.jvector.vector.ByteVectorSimilarityFunction; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import org.yaml.snakeyaml.Yaml; @@ -124,6 +125,18 @@ default LoadBehavior loadBehavior() { return LoadBehavior.LEGACY_SCRUB; } + /** + * Maps the dataset's similarity function to its byte-vector equivalent by name. + * Both enums share the same three value names (EUCLIDEAN, DOT_PRODUCT, COSINE), + * so this is a direct valueOf mapping — no new YAML key is required. + * + * @return the byte similarity function, or empty if no similarity function is configured + */ + default Optional byteSimilarityFunction() { + return similarityFunction() + .map(vsf -> ByteVectorSimilarityFunction.valueOf(vsf.name())); + } + /** * A convenience method to capture the notion of a valid dataset. * As any additional qualifiers are added to this data carrier, this method should be updated accordingly. diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetUtils.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetUtils.java index 61dc64652..8ce3dde81 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetUtils.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetUtils.java @@ -43,18 +43,18 @@ public class DataSetUtils { /** * Processes a dataset using the configured load behavior from the dataset metadata. */ - public static DataSet processDataSet(String pathStr, - DataSetProperties props, - List> baseVectors, - List> queryVectors, - List> groundTruth) { + public static DataSet processDataSet(String pathStr, + DataSetProperties props, + List> baseVectors, + List> queryVectors, + List> groundTruth) { var vsf = props.similarityFunction() .orElseThrow(() -> new IllegalArgumentException( "No similarity function configured for dataset: " + props.getName())); switch (props.loadBehavior()) { case NO_SCRUB: - return new SimpleDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); + return new FloatDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); case LEGACY_SCRUB: return legacyScrubDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); default: @@ -68,15 +68,15 @@ public static DataSet processDataSet(String pathStr, * so that load behavior is controlled explicitly by dataset metadata. */ @Deprecated(forRemoval = true) - public static DataSet getScrubbedDataSet(String pathStr, - VectorSimilarityFunction vsf, - List> baseVectors, - List> queryVectors, - List> groundTruth) { + public static DataSet getScrubbedDataSet(String pathStr, + VectorSimilarityFunction vsf, + List> baseVectors, + List> queryVectors, + List> groundTruth) { return legacyScrubDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); } - private static DataSet legacyScrubDataSet(String pathStr, + private static DataSet legacyScrubDataSet(String pathStr, VectorSimilarityFunction vsf, List> baseVectors, List> queryVectors, @@ -119,7 +119,7 @@ private static DataSet legacyScrubDataSet(String pathStr, } assert scrubbedQueryVectors.size() == gtSet.size(); - return new SimpleDataSet(pathStr, vsf, scrubbedBaseVectors, scrubbedQueryVectors, gtSet); + return new FloatDataSet(pathStr, vsf, scrubbedBaseVectors, scrubbedQueryVectors, gtSet); } private static boolean isValidLegacyVector(VectorFloat vector, VectorSimilarityFunction vsf) { diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/SimpleDataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/FloatDataSet.java similarity index 79% rename from jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/SimpleDataSet.java rename to jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/FloatDataSet.java index bf9c69376..2500c1817 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/SimpleDataSet.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/FloatDataSet.java @@ -23,7 +23,13 @@ import java.util.List; -public class SimpleDataSet implements DataSet { +/** + * A float32 dataset backed by in-memory {@link VectorFloat} lists. + * + *

Type-specific accessors ({@link #getBaseRavv()}, {@link #getSimilarityFunction()}) live here + * rather than on the {@link DataSet} interface. + */ +public class FloatDataSet implements DataSet> { private final String name; private final VectorSimilarityFunction similarityFunction; private final List> baseVectors; @@ -31,11 +37,11 @@ public class SimpleDataSet implements DataSet { private final List> groundTruth; private RandomAccessVectorValues baseRavv; - public SimpleDataSet(String name, - VectorSimilarityFunction similarityFunction, - List> baseVectors, - List> queryVectors, - List> groundTruth) + public FloatDataSet(String name, + VectorSimilarityFunction similarityFunction, + List> baseVectors, + List> queryVectors, + List> groundTruth) { if (baseVectors.isEmpty()) { throw new IllegalArgumentException("Base vectors must not be empty"); @@ -44,9 +50,8 @@ public SimpleDataSet(String name, throw new IllegalArgumentException("Query vectors must not be empty"); } if (groundTruth.isEmpty()) { - throw new IllegalArgumentException("Ground truth vectors must not be empty"); + throw new IllegalArgumentException("Ground truth must not be empty"); } - if (baseVectors.get(0).length() != queryVectors.get(0).length()) { throw new IllegalArgumentException("Base and query vectors must have the same dimensionality"); } @@ -64,27 +69,14 @@ public SimpleDataSet(String name, name, baseVectors.size(), queryVectors.size(), baseVectors.get(0).length()); } - @Override - public int getDimension() { - return getBaseVectors().get(0).length(); - } - - @Override - public RandomAccessVectorValues getBaseRavv() { - if (baseRavv == null) { - baseRavv = new ListRandomAccessVectorValues(getBaseVectors(), getDimension()); - } - return baseRavv; - } - @Override public String getName() { return name; } @Override - public VectorSimilarityFunction getSimilarityFunction() { - return similarityFunction; + public int getDimension() { + return baseVectors.get(0).length(); } @Override @@ -101,4 +93,15 @@ public List> getQueryVectors() { public List> getGroundTruth() { return groundTruth; } + + public RandomAccessVectorValues getBaseRavv() { + if (baseRavv == null) { + baseRavv = new ListRandomAccessVectorValues(baseVectors, getDimension()); + } + return baseRavv; + } + + public VectorSimilarityFunction getSimilarityFunction() { + return similarityFunction; + } } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/DatasetInfoWriter.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/DatasetInfoWriter.java index fd3b94ce3..a1d19c065 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/DatasetInfoWriter.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/DatasetInfoWriter.java @@ -16,7 +16,7 @@ package io.github.jbellis.jvector.example.reporting; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import java.io.IOException; import java.nio.charset.StandardCharsets; @@ -114,7 +114,7 @@ public static Row fromDataSet(String datasetName, String basePath, String queryPath, String groundTruthPath, - DataSet ds) { + FloatDataSet ds) { return new Row( datasetName, basePath, diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/RunArtifacts.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/RunArtifacts.java index 5cb2ebcde..0be6eae7b 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/RunArtifacts.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/reporting/RunArtifacts.java @@ -17,6 +17,7 @@ package io.github.jbellis.jvector.example.reporting; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.Metric; import io.github.jbellis.jvector.example.yaml.MultiConfig; import io.github.jbellis.jvector.example.yaml.MetricSelection; @@ -236,11 +237,11 @@ public void logRow(String datasetName, public Map> benchmarksToLog() { return benchmarksToLog; } public MetricSelection metricsToLog() { return metricsToLog; } - public void registerDataset(String datasetName, DataSet ds) throws IOException { + public void registerDataset(String datasetName, DataSet ds) throws IOException { if (datasetInfoWriter == null) { return; // disabled } - datasetInfoWriter.register(DatasetInfoWriter.fromDataSet(datasetName, "", "", "", ds)); + datasetInfoWriter.register(DatasetInfoWriter.fromDataSet(datasetName, "", "", "", (FloatDataSet) ds)); } } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/DiskIntro.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/DiskIntro.java index cfb70da09..affc85cc8 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/DiskIntro.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/DiskIntro.java @@ -25,7 +25,7 @@ import io.github.jbellis.jvector.disk.ReaderSupplier; import io.github.jbellis.jvector.disk.ReaderSupplierFactory; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.example.util.AccuracyMetrics; import io.github.jbellis.jvector.graph.GraphIndexBuilder; @@ -50,7 +50,7 @@ public class DiskIntro { public static void main(String[] args) throws IOException { // This is a preconfigured dataset that will be downloaded automatically. - DataSet dataset = DataSets.loadDataSet("ada002-100k").orElseThrow(() -> + FloatDataSet dataset = (FloatDataSet) DataSets.loadDataSet("ada002-100k").orElseThrow(() -> new RuntimeException("Dataset doesn't exist or wasn't configured correctly") ).getDataSet(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/LargerThanMemory.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/LargerThanMemory.java index 5f22b1cbc..7a1c1a8c1 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/LargerThanMemory.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/LargerThanMemory.java @@ -29,7 +29,7 @@ import io.github.jbellis.jvector.disk.RandomAccessReader; import io.github.jbellis.jvector.disk.ReaderSupplier; import io.github.jbellis.jvector.disk.ReaderSupplierFactory; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.example.util.AccuracyMetrics; import io.github.jbellis.jvector.graph.GraphIndexBuilder; @@ -60,7 +60,7 @@ public static void main(String[] args) throws IOException { // The DataSet provided by loadDataSet is in-memory, // but you can apply the same technique even when you don't have // the base vectors in-memory. - DataSet dataset = DataSets.loadDataSet("e5-small-v2-100k").orElseThrow(() -> + FloatDataSet dataset = (FloatDataSet) DataSets.loadDataSet("e5-small-v2-100k").orElseThrow(() -> new RuntimeException("Dataset doesn't exist or wasn't configured correctly") ).getDataSet(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/NvqExample.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/NvqExample.java index 2f0e9b1eb..58a7161c5 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/NvqExample.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/tutorial/NvqExample.java @@ -27,6 +27,7 @@ import io.github.jbellis.jvector.disk.ReaderSupplierFactory; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.util.AccuracyMetrics; import io.github.jbellis.jvector.graph.GraphIndexBuilder; import io.github.jbellis.jvector.graph.GraphSearcher; @@ -52,7 +53,7 @@ public class NvqExample { public static void main(String[] args) throws IOException { // Load a preconfigured dataset - var ds = DataSets.loadDataSet("ada002-100k").orElseThrow(() -> + var ds = (FloatDataSet) DataSets.loadDataSet("ada002-100k").orElseThrow(() -> new RuntimeException("dataset not found")) .getDataSet(); var dim = ds.getDimension(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/CompressorParameters.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/CompressorParameters.java index 2f4aceaf7..40eafd70a 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/CompressorParameters.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/CompressorParameters.java @@ -16,7 +16,7 @@ package io.github.jbellis.jvector.example.util; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.quantization.BinaryQuantization; import io.github.jbellis.jvector.quantization.NVQuantization; import io.github.jbellis.jvector.quantization.ProductQuantization; @@ -29,12 +29,12 @@ public boolean supportsCaching() { return false; } - public String idStringFor(DataSet ds) { + public String idStringFor(FloatDataSet ds) { // only required when supportsCaching() is true throw new UnsupportedOperationException(); } - public abstract VectorCompressor computeCompressor(DataSet ds); + public abstract VectorCompressor computeCompressor(FloatDataSet ds); public static class PQParameters extends CompressorParameters { private final int m; @@ -50,12 +50,12 @@ public PQParameters(int m, int k, boolean isCentered, float anisotropicThreshold } @Override - public VectorCompressor computeCompressor(DataSet ds) { + public VectorCompressor computeCompressor(FloatDataSet ds) { return ProductQuantization.compute(ds.getBaseRavv(), m, k, isCentered, anisotropicThreshold); } @Override - public String idStringFor(DataSet ds) { + public String idStringFor(FloatDataSet ds) { return String.format("PQ_%s_%d_%d_%s_%s", ds.getName(), m, k, isCentered, anisotropicThreshold); } @@ -67,7 +67,7 @@ public boolean supportsCaching() { public static class BQParameters extends CompressorParameters { @Override - public VectorCompressor computeCompressor(DataSet ds) { + public VectorCompressor computeCompressor(FloatDataSet ds) { return new BinaryQuantization(ds.getDimension()); } } @@ -80,12 +80,12 @@ public NVQParameters(int nSubVectors) { } @Override - public VectorCompressor computeCompressor(DataSet ds) { + public VectorCompressor computeCompressor(FloatDataSet ds) { return NVQuantization.compute(ds.getBaseRavv(), nSubVectors); } @Override - public String idStringFor(DataSet ds) { + public String idStringFor(FloatDataSet ds) { return String.format("NVQ_%s_%d_%s", ds.getName(), nSubVectors); } @@ -97,7 +97,7 @@ public boolean supportsCaching() { private static class NoCompressionParameters extends CompressorParameters { @Override - public VectorCompressor computeCompressor(DataSet ds) { + public VectorCompressor computeCompressor(FloatDataSet ds) { return null; } } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/DataSetPartitioner.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/DataSetPartitioner.java index 1e6a83f40..d80946693 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/DataSetPartitioner.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/DataSetPartitioner.java @@ -16,7 +16,7 @@ package io.github.jbellis.jvector.example.util; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.yaml.TestDataPartition; import io.github.jbellis.jvector.vector.types.VectorFloat; @@ -39,7 +39,7 @@ public PartitionedData(List>> vectors, List sizes) } } - public static PartitionedData partition(DataSet ds, int numParts, TestDataPartition.Distribution distribution) { + public static PartitionedData partition(FloatDataSet ds, int numParts, TestDataPartition.Distribution distribution) { return partition(ds.getBaseVectors(), numParts, distribution); } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/SiftLoader.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/SiftLoader.java index a491d0c9e..6a90fd89d 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/SiftLoader.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/SiftLoader.java @@ -17,6 +17,7 @@ package io.github.jbellis.jvector.example.util; import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.ByteSequence; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; @@ -60,6 +61,31 @@ public static List> readFvecs(String filePath) { return vectors; } + public static List> readBvecs(String filePath) { + var vectors = new ArrayList>(); + try (var dis = new DataInputStream(new BufferedInputStream(new FileInputStream(filePath)))) { + while (dis.available() > 0) { + var dimension = Integer.reverseBytes(dis.readInt()); + if (dimension <= 0) { + throw new IOException("Corrupt bvecs file: negative or zero dimension " + dimension + " (possible file corruption or wrong format)"); + } + if (dimension > 100_000) { + throw new IOException("Unreasonable dimension " + dimension + " in bvecs file (possible file corruption or wrong format)"); + } + var buffer = new byte[dimension]; + dis.readFully(buffer); + var bv = vectorTypeSupport.createByteSequence(dimension); + for (int i = 0; i < dimension; i++) { + bv.set(i, buffer[i]); + } + vectors.add(bv); + } + } catch (IOException ex) { + throw new UncheckedIOException(ex); + } + return vectors; + } + public static List> readIvecs(String filename) { var groundTruthTopK = new ArrayList>(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/CommonParameters.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/CommonParameters.java index 626efc591..444f55daa 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/CommonParameters.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/CommonParameters.java @@ -16,8 +16,10 @@ package io.github.jbellis.jvector.example.yaml; -import io.github.jbellis.jvector.example.util.CompressorParameters; +import io.github.jbellis.jvector.example.benchmarks.datasets.ByteDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.util.CompressorParameters; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import java.util.List; import java.util.function.Function; @@ -26,7 +28,23 @@ public class CommonParameters { public List compression; - public List> getCompressorParameters() { + public List> getCompressorParameters(DataSet ds) { + if (ds instanceof ByteDataSet) { + if (compression != null) { + for (var c : compression) { + if (c.type != null && !c.type.equalsIgnoreCase("None")) { + throw new IllegalArgumentException(String.format( + "Compression type '%s' is not supported for INT8 dataset '%s'. INT8 datasets do not support compression.", + c.type, ds.getName())); + } + } + } + return List.of(__ -> CompressorParameters.NONE); + } + + if (compression == null) { + return List.of(__ -> CompressorParameters.NONE); + } return compression.stream().map(Compression::getCompressorParameters).collect(Collectors.toList()); } } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/Compression.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/Compression.java index a8277508a..aa0583a04 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/Compression.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/Compression.java @@ -17,7 +17,7 @@ package io.github.jbellis.jvector.example.yaml; import io.github.jbellis.jvector.example.util.CompressorParameters; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; import java.util.Map; @@ -27,7 +27,7 @@ public class Compression { public String type; public Map parameters; - public Function getCompressorParameters() { + public Function getCompressorParameters() { switch (type) { case "None": return __ -> CompressorParameters.NONE; diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/ConstructionParameters.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/ConstructionParameters.java index 5177fdf4a..cd592c8eb 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/ConstructionParameters.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/ConstructionParameters.java @@ -16,6 +16,8 @@ package io.github.jbellis.jvector.example.yaml; +import io.github.jbellis.jvector.example.benchmarks.datasets.ByteDataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; import java.util.EnumSet; @@ -33,40 +35,52 @@ public class ConstructionParameters extends CommonParameters { public List fusedGraph; public Boolean useSavedIndexIfExists; - public List> getFeatureSets() { + public List> getFeatureSets(DataSet ds) { + if (ds instanceof ByteDataSet) { + if (reranking != null && !reranking.isEmpty()) { + throw new IllegalArgumentException(String.format( + "Reranking %s is not supported for INT8 dataset '%s'. INT8 datasets do not support reranking.", + reranking, ds.getName())); + } + return List.of(EnumSet.of(FeatureId.INLINE_BYTE_VECTORS)); + } + List> featureSets = null; - for (var fusedItem : fusedGraph) { - var newFeatures = reranking.stream().map(item -> { - EnumSet features; + if (fusedGraph != null && reranking != null) { + for (var fusedItem : fusedGraph) { + var newFeatures = reranking.stream().map(item -> { + EnumSet features; - switch (item) { - case "FP": - if (fusedItem) { - features = EnumSet.of(FeatureId.INLINE_VECTORS, FeatureId.FUSED_PQ); - } else { - features = EnumSet.of(FeatureId.INLINE_VECTORS); - } - break; - case "NVQ": - if (fusedItem) { - features = EnumSet.of(FeatureId.NVQ_VECTORS, FeatureId.FUSED_PQ); - } else { - features = EnumSet.of(FeatureId.NVQ_VECTORS); - } - break; - default: - throw new IllegalArgumentException("Only 'FP' and 'NVQ' are supported"); - } + switch (item) { + case "FP": + if (fusedItem) { + features = EnumSet.of(FeatureId.INLINE_VECTORS, FeatureId.FUSED_PQ); + } else { + features = EnumSet.of(FeatureId.INLINE_VECTORS); + } + break; + case "NVQ": + if (fusedItem) { + features = EnumSet.of(FeatureId.NVQ_VECTORS, FeatureId.FUSED_PQ); + } else { + features = EnumSet.of(FeatureId.NVQ_VECTORS); + } + break; + default: + throw new IllegalArgumentException("Only 'FP' and 'NVQ' are supported"); + } - return features; - }).collect(Collectors.toList()); - if (featureSets == null) { - featureSets = newFeatures; - } else { - featureSets.addAll(newFeatures); + return features; + }).collect(Collectors.toList()); + if (featureSets == null) { + featureSets = newFeatures; + } else { + featureSets.addAll(newFeatures); + } } } return featureSets; } + } \ No newline at end of file diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/graph/disk/ParallelWriteExample.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/graph/disk/ParallelWriteExample.java index f3728234c..dfd4759fd 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/graph/disk/ParallelWriteExample.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/graph/disk/ParallelWriteExample.java @@ -17,7 +17,7 @@ package io.github.jbellis.jvector.graph.disk; import io.github.jbellis.jvector.disk.ReaderSupplierFactory; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.graph.GraphIndexBuilder; import io.github.jbellis.jvector.graph.ImmutableGraphIndex; @@ -302,7 +302,7 @@ public static void main(String[] args) throws IOException { String datasetName = args.length > 0 ? args[0] : "cohere-english-v3-100k"; System.out.println("Loading dataset: " + datasetName); - DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( + FloatDataSet ds = (FloatDataSet) DataSets.loadDataSet(datasetName).orElseThrow( () -> new RuntimeException("Dataset " + datasetName + " not found") ).getDataSet(); System.out.printf("Loaded %d vectors of dimension %d%n", ds.getBaseVectors().size(), ds.getDimension()); diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetPropertiesTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetPropertiesTest.java index 13f2136aa..d76a2da6a 100644 --- a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetPropertiesTest.java +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetPropertiesTest.java @@ -345,11 +345,9 @@ public void dataSetInfoLazyLoading() { var callCount = new int[]{0}; var base = new DataSetProperties.PropertyMap(Map.of(DataSetProperties.KEY_NAME, "lazy")); // Return a dummy non-null sentinel so the cache works (null would defeat the null-check) - var sentinel = new DataSet() { + var sentinel = new DataSet>() { public int getDimension() { return 0; } - public RandomAccessVectorValues getBaseRavv() { return null; } public String getName() { return "sentinel"; } - public VectorSimilarityFunction getSimilarityFunction() { return VectorSimilarityFunction.COSINE; } public List> getBaseVectors() { return Collections.emptyList(); } public List> getQueryVectors() { return Collections.emptyList(); } public List> getGroundTruth() { return Collections.emptyList(); } diff --git a/jvector-examples/yaml-configs/dataset-catalogs/local-catalog.yaml b/jvector-examples/yaml-configs/dataset-catalogs/local-catalog.yaml index 0aee43b9a..b7f37bf77 100644 --- a/jvector-examples/yaml-configs/dataset-catalogs/local-catalog.yaml +++ b/jvector-examples/yaml-configs/dataset-catalogs/local-catalog.yaml @@ -39,3 +39,8 @@ # base: private_base.fvecs # query: private_query.fvecs # gt: private_gt.ivecs + +sift1m-128-euclidean-int8: + base: /home/raghuveer/myplayground/dataset-processing/runs/sift1m_hdf5/sift1m_base_975462.bvecs + query: /home/raghuveer/myplayground/dataset-processing/runs/sift1m_hdf5/sift1m_query_10000.bvecs + gt: /home/raghuveer/myplayground/dataset-processing/runs/sift1m_hdf5/sift1m_gt_l2_100.ivecs diff --git a/jvector-examples/yaml-configs/dataset-metadata.yml b/jvector-examples/yaml-configs/dataset-metadata.yml index a0c840a8c..25ab66824 100644 --- a/jvector-examples/yaml-configs/dataset-metadata.yml +++ b/jvector-examples/yaml-configs/dataset-metadata.yml @@ -108,4 +108,7 @@ openai-3072-1m: load_behavior: NO_SCRUB openai-1536-1m: similarity_function: DOT_PRODUCT + load_behavior: NO_SCRUB +sift1m-128-euclidean-int8: + similarity_function: EUCLIDEAN load_behavior: NO_SCRUB \ No newline at end of file diff --git a/jvector-examples/yaml-configs/index-parameters/sift1m-128-euclidean-int8.yml b/jvector-examples/yaml-configs/index-parameters/sift1m-128-euclidean-int8.yml new file mode 100644 index 000000000..0abd4cff9 --- /dev/null +++ b/jvector-examples/yaml-configs/index-parameters/sift1m-128-euclidean-int8.yml @@ -0,0 +1,23 @@ +yamlSchemaVersion: 1 +onDiskIndexVersion: 6 + +dataset: sift1m-128-euclidean-int8 + +construction: + outDegree: [32] + efConstruction: [100] + neighborOverflow: [1.2f] + addHierarchy: [No] + refineFinalGraph: [No] + fusedGraph: [No] + compression: + - type: None + useSavedIndexIfExists: No + +search: + topKOverquery: + 10: [1.0] + 100: [1.0] + useSearchPruning: [Yes] + compression: + - type: None diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/microbench/GraphBuildBench.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/microbench/GraphBuildBench.java index 28127fb34..33875b3de 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/microbench/GraphBuildBench.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/microbench/GraphBuildBench.java @@ -16,8 +16,8 @@ package io.github.jbellis.jvector.microbench; -import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; +import io.github.jbellis.jvector.example.benchmarks.datasets.FloatDataSet; import io.github.jbellis.jvector.graph.GraphIndexBuilder; import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; import org.openjdk.jmh.annotations.Benchmark; @@ -40,11 +40,11 @@ public class GraphBuildBench { @State(Scope.Benchmark) public static class Parameters { - final DataSet ds; + final FloatDataSet ds; final ListRandomAccessVectorValues ravv; public Parameters() { - this.ds = DataSets.loadDataSet("glove-100-angular").orElseThrow( + this.ds = (FloatDataSet) DataSets.loadDataSet("glove-100-angular").orElseThrow( () -> new RuntimeException("Unable to load dataset: glove-100-angular") ).getDataSet(); this.ravv = new ListRandomAccessVectorValues(ds.getBaseVectors(), ds.getBaseVectors().get(0).length());