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/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-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/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-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-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..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 @@ -16,15 +16,20 @@ 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.graph.VectorValues; +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; /** * Encapsulates comparing node distances for GraphIndexBuilder. @@ -59,6 +64,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. *

@@ -106,8 +125,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 = asRandomAccessSupplier(vectorsRaw); + Supplier vectorsCopy = asRandomAccessSupplier(vectorsCopyRaw); return new BuildScoreProvider() { @Override @@ -211,6 +232,78 @@ 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}. + * 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 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..33d2875ef --- /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 * 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() * (128.0f * 128.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); +} 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-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/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/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/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]); } 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-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/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 52bdc872a..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; @@ -433,7 +436,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, () -> { @@ -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 c76796bb4..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; @@ -340,7 +348,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)); @@ -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/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()); 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..8a3e08e8f --- /dev/null +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/vector/TestByteVectorSimilarityFunction.java @@ -0,0 +1,186 @@ +/* + * 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 * 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); + } + + @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 -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); + } +} 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..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 @@ -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,121 @@ 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); + } + + /** + * Verifies that the SIMD byte-vector kernels agree with the scalar baseline for + * several dimensions that stress boundary conditions in the vectorised loops: + *

    + *
  • dim=1 — the absolute minimum case
  • + *
  • dim=8 — exact multiple of the smallest SIMD lane width
  • + *
  • dim=128 — exact multiple of larger lane widths
  • + *
  • dim=256 — exact multiple of AVX2/AVX512 lane widths
  • + *
  • dim=255 — one less than 256, exposes tail-loop handling
  • + *
+ */ + @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); 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..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 @@ -976,14 +976,250 @@ 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 dot product of two signed int8 byte vectors. + */ + @Override + public float dotProduct(ByteSequence a, ByteSequence b) { + 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 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + acc = acc.add(va.mul(vb)); + } + + int result = acc.reduceLanes(VectorOperators.ADD); + for (int i = limit; i < length; i++) { + result += a.get(i) * b.get(i); + } + return result; + } + + float dotProductBytes256(ByteSequence a, ByteSequence b) { + final int length = a.length(); + 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, 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(i) * b.get(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(); + int result = 0; + for (int i = 0; i < length; i++) { + result += a.get(i) * b.get(i); + } + return result; + } + /** - * 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 sum of squared differences between two signed int8 byte vectors. */ + @Override + public float squareDistance(ByteSequence a, ByteSequence b) { + 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 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 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(i) - b.get(i); + result += diff * diff; + } + return result; + } + + float squareDistanceBytes256(ByteSequence a, ByteSequence b) { + final int length = a.length(); + 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, 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(i) - b.get(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(); + int result = 0; + for (int i = 0; i < length; i++) { + int diff = a.get(i) - b.get(i); + result += diff * diff; + } + return result; + } + + /** + * Vectorized cosine similarity between two signed int8 byte vectors. + */ + @Override + public float cosine(ByteSequence a, ByteSequence b) { + 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 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_512, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_128, b, 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(i), bi = b.get(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 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, i) + .convertShape(VectorOperators.B2I, IntVector.SPECIES_256, 0) + .reinterpretAsInts(); + IntVector vb = fromByteSequence(ByteVector.SPECIES_64, b, 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(i), bi = b.get(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(); + long dotResult = 0, normAResult = 0, normBResult = 0; + for (int i = 0; i < length; i++) { + int ai = a.get(i), bi = b.get(i); + dotResult += ai * bi; + normAResult += ai * ai; + normBResult += bi * bi; + } + return (float) (dotResult / Math.sqrt((double) normAResult * normBResult)); + } + @Override public int hammingDistance(long[] a, long[] b) { var sum = LongVector.zero(LongVector.SPECIES_PREFERRED);