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..1a7c3cf2e 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 @@ -249,7 +249,6 @@ private static void writeCompletedCount(int count) { // ---------- Benchmark state ---------- private RandomAccessVectorValues ravv; private List> queryVectors; - private List> baseVectors; private List> groundTruth; private DataSet ds; private VectorSimilarityFunction similarityFunction; @@ -399,7 +398,6 @@ public void setup() throws Exception { if (datasetPortion == 1.0) { ravv = ds.getBaseRavv(); - baseVectors = ds.getBaseVectors(); } else { int totalVectors = ds.getBaseRavv().size(); int portionedSize = (int) (totalVectors * datasetPortion); @@ -408,8 +406,7 @@ public void setup() throws Exception { "datasetPortion=" + datasetPortion + " yields " + portionedSize + " vectors, fewer than numPartitions=" + numPartitions); } - baseVectors = ds.getBaseVectors().subList(0, portionedSize); - ravv = new ListRandomAccessVectorValues(baseVectors, ds.getDimension()); + ravv = ds.getBaseRavv().range(0, portionedSize); } similarityFunction = ds.getSimilarityFunction(); @@ -428,7 +425,6 @@ public void setup() throws Exception { } } else { ravv = null; - baseVectors = null; dimension = -1; if (needsRecallData) { @@ -469,7 +465,7 @@ public void setup() throws Exception { } if (workloadMode == WorkloadMode.PARTITION || workloadMode == WorkloadMode.PARTITION_AND_COMPACT) { - var partitionedData = DataSetPartitioner.partition(baseVectors, numPartitions, splitDistribution); + var partitionedData = DataSetPartitioner.partition(ravv, numPartitions, splitDistribution); vectorsPerSourceCount = partitionedData.sizes; } else { vectorsPerSourceCount = null; @@ -479,7 +475,7 @@ public void setup() throws Exception { if (jfrPartitioning) { jfrPartitioningRecorder.start(JFR_DIR, "partitioning-" + jfrParamSuffix() + ".jfr", jfrObjectCount); } - buildPartitions(ds, baseVectors); + buildPartitions(ravv); if (jfrPartitioningRecorder.isActive()) { jfrPartitioningRecorder.stop(); } @@ -595,7 +591,7 @@ private void verifyPartitionsExist(Path partitionsDir, int numPartitions) { } } - private void buildPartitions(DataSet ds, List> baseVectors) throws Exception { + private void buildPartitions(RandomAccessVectorValues baseVectors) throws Exception { var partitionedData = DataSetPartitioner.partition(baseVectors, numPartitions, splitDistribution); vectorsPerSourceCount = partitionedData.sizes; @@ -604,9 +600,9 @@ private void buildPartitions(DataSet ds, List> baseVectors) throw numPartitions, partitionsBaseDir.toAbsolutePath(), graphDegree, beamWidth, splitDistribution, vectorsPerSourceCount, indexPrecision, parallelWriteThreads, resolvedVectorizationProvider); - int dimension = baseVectors.get(0).length(); + int dimension = baseVectors.dimension(); for (int i = 0; i < numPartitions; i++) { - List> vectorsPerSource = partitionedData.vectors.get(i); + RandomAccessVectorValues ravvPerSource = partitionedData.vectors.get(i); // Round-robin assignment of partition files to storage paths, but still keep canonical base dir name stable. Path baseDirForThisSegment = storagePaths.get(i % storagePaths.size()); @@ -616,9 +612,8 @@ private void buildPartitions(DataSet ds, List> baseVectors) throw } log.info("Building partition {}/{}: vectors={} -> {}", - i + 1, numPartitions, vectorsPerSource.size(), outputPath.toAbsolutePath()); + i + 1, numPartitions, ravvPerSource.size(), outputPath.toAbsolutePath()); - var ravvPerSource = new ListRandomAccessVectorValues(vectorsPerSource, dimension); BuildScoreProvider bspPerSource; ProductQuantization pq = null; PQVectors pqVectors = null; @@ -655,7 +650,8 @@ private void buildPartitions(DataSet ds, List> baseVectors) throw try (var writer = writerBuilder.build()) { var suppliers = new EnumMap>(FeatureId.class); - suppliers.put(FeatureId.INLINE_VECTORS, ordinal -> new InlineVectors.State(ravvPerSource.getVector(ordinal))); + var vectors = ravvPerSource.threadLocalSupplier(); // the parallel writer calls suppliers from worker threads + suppliers.put(FeatureId.INLINE_VECTORS, ordinal -> new InlineVectors.State(vectors.get().getVector(ordinal))); if (indexPrecision == IndexPrecision.FUSEDPQ) { var view = graph.getView(); @@ -713,7 +709,7 @@ private long compactPartitions() throws Exception { return compactionTimeMs; } - private long buildFromScratch(List> baseVectors) throws Exception { + private long buildFromScratch(RandomAccessVectorValues baseVectors) throws Exception { if (scratchOutputPath.getParent() != null) { Files.createDirectories(scratchOutputPath.getParent()); } @@ -721,8 +717,8 @@ private long buildFromScratch(List> baseVectors) throws Exception Files.delete(scratchOutputPath); } - int dimension = baseVectors.get(0).length(); - var full = new ListRandomAccessVectorValues(baseVectors, dimension); + int dimension = baseVectors.dimension(); + var full = baseVectors; log.info("Building from scratch: vectors={} dim={} sim={} deg={} bw={} precision={} pwThreads={} vp={} -> {}", full.size(), dimension, similarityFunction, @@ -761,7 +757,8 @@ private long buildFromScratch(List> baseVectors) throws Exception try (var writer = writerBuilder.build()) { var suppliers = new EnumMap>(FeatureId.class); - suppliers.put(FeatureId.INLINE_VECTORS, ord -> new InlineVectors.State(full.getVector(ord))); + var vectors = full.threadLocalSupplier(); // the parallel writer calls suppliers from worker threads + suppliers.put(FeatureId.INLINE_VECTORS, ord -> new InlineVectors.State(vectors.get().getVector(ord))); if (indexPrecision == IndexPrecision.FUSEDPQ) { var view = graph.getView(); @@ -843,7 +840,7 @@ public void run(Blackhole blackhole, RecallResult recallResult) throws Exception break; case BUILD: - durationMs = buildFromScratch(baseVectors); + durationMs = buildFromScratch(ravv); if (measureRecall) { searchStats = runRecall(scratchOutputPath); recall = searchStats.recall; diff --git a/docs/benchmarking.md b/docs/benchmarking.md index 79ddba82e..74ae8e8b2 100644 --- a/docs/benchmarking.md +++ b/docs/benchmarking.md @@ -33,6 +33,21 @@ Datasets are grouped into categories. The categories can be arbitrarily chosen f Dataset similarity functions are configured in `jvector-examples/yaml-configs/dataset-metadata.yml`. +Each entry may carry a loader *profile* and a list of *wrappers* that change how the base vectors are held after loading. Two forms are accepted: + +```yaml +regression-tests: + - cap-1M # name only: default profile, base vectors cached in heap memory + - cohere-english-v3-1M(mmap) # sugared: name, then wrappers in parentheses + - name: cohere-english-v3-10M # structured + profile: default # optional; "default" when omitted + wrappers: # optional; applied left to right + - mmap + - lru: { grain: 4096, capacityMb: 512 } # a wrapper with options +``` + +The sugared form is `name`, `name:profile`, `name(wrapper,...)` or `name:profile(wrapper,...)`, where a wrapper may carry options in brackets as in `lru[grain=4096,capacityMb=512]`, and is also accepted anywhere a dataset name is (command-line patterns, `dataset:` in an index-parameters file). Built-in wrappers are `memory` (cache the base vectors in heap memory, the default when no wrappers are given), `mmap` (serve them from a memory-mapped fvecs file) and `lru` (a bounded least-recently-used cache of vector grains in front of the mapped file, for datasets larger than memory; options `grain` (vectors per grain) and `capacityMb`, defaulting to `-Djvector.dataset.lru.grain` and `-Djvector.dataset.lru.capacityMb` or 1024 and a quarter of the heap). Loaders that do not understand profiles accept `default` and reject any other profile. + Example `datasets.yml`: ```yaml 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..3444d06a7 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 @@ -78,6 +78,27 @@ default void getVectorInto(int node, VectorFloat destinationVector, int offse destinationVector.copyFrom(getVector(node), 0, offset, dimension()); } + /** + * Returns a view of the ordinal range {@code [fromOrdinal, toOrdinal)} of this RAVV, re-based so that + * ordinal {@code i} of the view reads ordinal {@code fromOrdinal + i} of this RAVV. + *

+ * The view does not copy vectors: it shares the underlying storage and inherits the sharing semantics + * of {@link #isValueShared()}. Its size is fixed at {@code toOrdinal - fromOrdinal} when this method is + * called, even if this RAVV later grows. Passing {@code 0} and {@link #size()} yields a view equivalent + * to this RAVV. + *

+ * The default implementation returns a {@link RangeRandomAccessVectorValues}; implementations with a cheaper + * native ranged read may override this. + * + * @param fromOrdinal the first ordinal of the range, inclusive; must be ≥ 0 + * @param toOrdinal the last ordinal of the range, exclusive; must be ≥ {@code fromOrdinal} and ≤ {@link #size()} + * @return a RAVV of size {@code toOrdinal - fromOrdinal} over the requested range + * @throws IndexOutOfBoundsException if the range is not within {@code [0, size()]} + */ + default RandomAccessVectorValues range(int fromOrdinal, int toOrdinal) { + return new RangeRandomAccessVectorValues(this, fromOrdinal, toOrdinal); + } + /** * @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. diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RangeRandomAccessVectorValues.java b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RangeRandomAccessVectorValues.java new file mode 100644 index 000000000..e9972d524 --- /dev/null +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/graph/RangeRandomAccessVectorValues.java @@ -0,0 +1,113 @@ +/* + * 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.VectorFloat; + +import java.util.Objects; + +/** + * A re-based view over a contiguous ordinal range {@code [fromOrdinal, toOrdinal)} of a backing + * {@link RandomAccessVectorValues}. Ordinal {@code i} of the view reads ordinal {@code fromOrdinal + i} + * of the backing RAVV. + *

+ * The view does not copy vectors. Its size is fixed at construction time, so a backing RAVV that grows + * afterwards (for example a {@link ListRandomAccessVectorValues} over a list that is still being appended to) + * is not reflected in the view. Value-sharing semantics are inherited from the backing RAVV. + *

+ * This is the default implementation returned by {@link RandomAccessVectorValues#range(int, int)}; implementations + * with a cheaper native ranged read may override that method instead. + */ +public final class RangeRandomAccessVectorValues implements RandomAccessVectorValues { + private final RandomAccessVectorValues ravv; + private final int fromOrdinal; + private final int size; + + /** + * Creates a view over ordinals {@code [fromOrdinal, toOrdinal)} of {@code ravv}. + * + * @param ravv the backing RAVV + * @param fromOrdinal the first backing ordinal of the range, inclusive + * @param toOrdinal the last backing ordinal of the range, exclusive + * @throws IndexOutOfBoundsException if the range is not within {@code [0, ravv.size()]} + */ + public RangeRandomAccessVectorValues(RandomAccessVectorValues ravv, int fromOrdinal, int toOrdinal) { + Objects.checkFromToIndex(fromOrdinal, toOrdinal, ravv.size()); + this.ravv = ravv; + this.fromOrdinal = fromOrdinal; + this.size = toOrdinal - fromOrdinal; + } + + /** + * @return the backing ordinal that view ordinal {@code 0} maps to + */ + public int fromOrdinal() { + return fromOrdinal; + } + + /** + * @return the exclusive upper bound of the backing ordinal range + */ + public int toOrdinal() { + return fromOrdinal + size; + } + + @Override + public int size() { + return size; + } + + @Override + public int dimension() { + return ravv.dimension(); + } + + @Override + public VectorFloat getVector(int nodeId) { + return ravv.getVector(fromOrdinal + Objects.checkIndex(nodeId, size)); + } + + @Override + public void getVectorInto(int node, VectorFloat destinationVector, int offset) { + ravv.getVectorInto(fromOrdinal + Objects.checkIndex(node, size), destinationVector, offset); + } + + @Override + public boolean isValueShared() { + return ravv.isValueShared(); + } + + /** + * Copies the backing RAVV and wraps the copy in an equivalent view. If the backing RAVV is un-shared + * and returns itself from {@link RandomAccessVectorValues#copy()}, this view returns itself as well. + */ + @Override + public RandomAccessVectorValues copy() { + RandomAccessVectorValues copied = ravv.copy(); + return copied == ravv ? this : new RangeRandomAccessVectorValues(copied, fromOrdinal, toOrdinal()); + } + + /** + * Narrows this view without adding a level of indirection: the result reads the backing RAVV directly + * at {@code [this.fromOrdinal + fromOrdinal, this.fromOrdinal + toOrdinal)}. + */ + @Override + public RandomAccessVectorValues range(int fromOrdinal, int toOrdinal) { + Objects.checkFromToIndex(fromOrdinal, toOrdinal, size); + return new RangeRandomAccessVectorValues(ravv, this.fromOrdinal + fromOrdinal, this.fromOrdinal + toOrdinal); + } +} 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..836eaf84c 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 @@ -126,7 +126,7 @@ public static void main(String[] args) throws IOException { DataSet ds = DataSets.loadDataSet(datasetName).orElseThrow( () -> new RuntimeException("Dataset " + datasetName + " not found") ).getDataSet(); - logger.info("Dataset loaded: {} with {} vectors", datasetName, ds.getBaseVectors().size()); + logger.info("Dataset loaded: {} with {} vectors", datasetName, ds.getBaseRavv().size()); String normalizedDatasetName = datasetName; if (normalizedDatasetName.endsWith(".hdf5")) { 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..552d5ff6a 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 @@ -23,7 +23,7 @@ import io.github.jbellis.jvector.example.util.CompactionPartitionSource; import io.github.jbellis.jvector.example.yaml.TestDataPartition.Distribution; import io.github.jbellis.jvector.graph.GraphSearcher; -import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import io.github.jbellis.jvector.graph.SearchResult; import io.github.jbellis.jvector.graph.disk.OnDiskGraphIndex; import io.github.jbellis.jvector.graph.disk.OnDiskGraphIndexCompactor; @@ -33,7 +33,6 @@ import io.github.jbellis.jvector.util.Bits; import io.github.jbellis.jvector.util.FixedBitSet; import io.github.jbellis.jvector.vector.VectorSimilarityFunction; -import io.github.jbellis.jvector.vector.types.VectorFloat; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -121,7 +120,7 @@ public static List run(DataSet ds) throws Exception { private static BenchResult runConfig(DataSet ds, PartitionConfig cfg) throws Exception { String datasetName = ds.getName(); logger.info("Compaction bench [{}] config {}: {} vectors", - datasetName, cfg.dirName(), ds.getBaseVectors().size()); + datasetName, cfg.dirName(), ds.getBaseRavv().size()); // 1. Fetch pre-built partitions from S3 (cached locally). List partitionPaths = CompactionPartitionSource.ensurePartitions( @@ -137,15 +136,14 @@ private static BenchResult runConfig(DataSet ds, PartitionConfig cfg) throws Exc private static BenchResult compactAndMeasure(DataSet ds, PartitionConfig cfg, List partitionPaths, Path tempDir) throws Exception { - List> baseVectors = ds.getBaseVectors(); - int dimension = ds.getDimension(); + RandomAccessVectorValues baseVectors = ds.getBaseRavv(); VectorSimilarityFunction vsf = ds.getSimilarityFunction(); String datasetName = ds.getName(); int numPartitions = cfg.numPartitions; // Load graphs and set up ordinal mapping: partition i's local ordinals shift by the sum of // all prior partition sizes, preserving the original base-vector ordering so global ordinal - // k maps back to baseVectors.get(k) by construction. + // k maps back to baseVectors ordinal k by construction. List rss = new ArrayList<>(numPartitions); List graphs = new ArrayList<>(numPartitions); List remappers = new ArrayList<>(numPartitions); @@ -184,8 +182,8 @@ private static BenchResult compactAndMeasure(DataSet ds, PartitionConfig cfg, rss.clear(); // Search the compacted graph: measure recall and search latency in one pass. - // Global ordinal k maps back to baseVectors.get(k) by construction. - SearchStats search = searchCompacted(compactPath, ds, baseVectors, dimension, vsf); + // Global ordinal k maps back to baseVectors ordinal k by construction. + SearchStats search = searchCompacted(compactPath, ds, baseVectors, vsf); logger.info(String.format( "%n" + " ┌─ Compaction result: %s [%s]%n" + @@ -244,11 +242,10 @@ static final class SearchStats { * mean and p99 per-query latency (ms) and throughput (queries/sec, single-threaded sequential). */ private static SearchStats searchCompacted(Path indexPath, DataSet ds, - List> baseVectors, - int dimension, VectorSimilarityFunction vsf) throws Exception { + RandomAccessVectorValues ravv, + VectorSimilarityFunction vsf) throws Exception { var queryVectors = ds.getQueryVectors(); var groundTruth = ds.getGroundTruth(); - var ravv = new ListRandomAccessVectorValues(baseVectors, dimension); try (var rs = ReaderSupplierFactory.open(indexPath)) { var graph = OnDiskGraphIndex.load(rs); 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..76ce96c68 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 @@ -331,7 +331,7 @@ static void runOneGraph(OnDiskGraphIndexCache cache, "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, ds.getBaseVectors().size(), (cv.ramBytesUsed() / 1024f / 1024f), encodingTimeS); + System.out.format("%s: %s encoded %d vectors [%.2f MB] in %.2fs%n", ds.getName(), compressor, ds.getBaseRavv().size(), (cv.ramBytesUsed() / 1024f / 1024f), encodingTimeS); } } @@ -506,13 +506,16 @@ private static BuilderWithSuppliers builderWithSuppliers(Set features var identityMapper = new OrdinalMapper.IdentityMapper(floatVectors.size() - 1); var builder = new RandomAccessOnDiskGraphIndexWriter.Builder(onHeapGraph, outPath); builder.withMapper(identityMapper); + // suppliers are invoked from several writer threads; a value-shared reader (e.g. a memory-mapped + // dataset) must be read through per-thread copies or the threads overwrite each other's vectors + var vectors = floatVectors.threadLocalSupplier(); Map> suppliers = new EnumMap<>(FeatureId.class); for (var featureId : features) { switch (featureId) { case INLINE_VECTORS: builder.with(new InlineVectors(floatVectors.dimension())); - suppliers.put(FeatureId.INLINE_VECTORS, ordinal -> new InlineVectors.State(floatVectors.getVector(ordinal))); + suppliers.put(FeatureId.INLINE_VECTORS, ordinal -> new InlineVectors.State(vectors.get().getVector(ordinal))); break; case FUSED_PQ: if (pq == null) { @@ -528,7 +531,7 @@ private static BuilderWithSuppliers builderWithSuppliers(Set features ? 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)))); + suppliers.put(FeatureId.NVQ_VECTORS, ordinal -> new NVQ.State(nvq.encode(vectors.get().getVector(ordinal)))); break; } @@ -885,7 +888,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, ds.getBaseRavv().size(), (cvArg.ramBytesUsed() / 1024f / 1024f)); } } 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..d46ae4fe1 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 @@ -24,6 +24,10 @@ /** * This provides a uniform way to access vector test data, regardless of where it comes from or how it is implemented. + *

+ * Base vectors are exposed only through {@link #getBaseRavv()}, so a dataset may keep them in heap memory, + * in a memory-mapped file, or anywhere else a {@link RandomAccessVectorValues} can read from. Callers that + * need a contiguous subset should use {@link RandomAccessVectorValues#range(int, int)} rather than copying. */ public interface DataSet { @@ -52,12 +56,6 @@ public interface DataSet { */ VectorSimilarityFunction getSimilarityFunction(); - /** - * The base vectors as a list. - * @return a list of base vectors - */ - List> getBaseVectors(); - /** * The query vectors as a list. * Each major index corresponds to the self-same index from {@link #getGroundTruth()}. @@ -69,7 +67,7 @@ public interface DataSet { /** * 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()}. + * Each minor index within represents the corresponding ordinal from {@link #getBaseRavv()}. * @return a list of query vectors. */ 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..eed30bfd9 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 @@ -114,6 +114,14 @@ public boolean isDuplicateVectorFree() { return baseProperties.isDuplicateVectorFree(); } + /** + * {@inheritDoc} + */ + @Override + public LoadBehavior loadBehavior() { + return baseProperties.loadBehavior(); + } + /// Returns the fully loaded and scrubbed {@link DataSet}. /// /// On the first invocation this triggers the deferred load pipeline, which may involve diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoader.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoader.java index 932ea2dc7..0fdddb48f 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoader.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoader.java @@ -47,4 +47,29 @@ public interface DataSetLoader { * @return a {@link DataSetInfo} handle for the dataset, if found */ Optional loadDataSet(String dataSetName); + + /** + * Looks up a dataset by {@link DataSetSpec}, honoring the spec's loader profile. + * + *

The default implementation is for loaders that do not understand profiles: it delegates to + * {@link #loadDataSet(String)} with the spec's name, accepts the {@value DataSetSpec#DEFAULT_PROFILE} + * profile silently, and throws if the dataset was found but a different profile was requested. A loader + * that does not recognise the name returns empty regardless of profile, so a later loader in the chain + * may still serve it. + * + *

Loaders that understand profiles override this method. Wrapper names on the spec are not the + * loader's concern; {@link DataSets} applies them to the loaded dataset. + * + * @param spec the dataset name and profile + * @return a {@link DataSetInfo} handle for the dataset, if found + * @throws IllegalArgumentException if this loader found the dataset but cannot honor a non-default profile + */ + default Optional loadDataSet(DataSetSpec spec) { + Optional found = loadDataSet(spec.getName()); + if (found.isPresent() && !spec.isDefaultProfile()) { + throw new IllegalArgumentException(getClass().getSimpleName() + " does not support dataset profiles, but profile '" + + spec.getProfile() + "' was requested for dataset '" + spec.getName() + "'"); + } + return found; + } } 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..838f5ee9c 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,7 +15,9 @@ */ package io.github.jbellis.jvector.example.benchmarks.datasets; +import io.github.jbellis.jvector.example.util.MappedFvecsRandomAccessVectorValues; import io.github.jbellis.jvector.example.util.SiftLoader; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.yaml.snakeyaml.Yaml; @@ -44,6 +46,7 @@ import java.security.NoSuchAlgorithmException; import java.util.HashMap; import java.util.LinkedHashMap; +import java.util.List; import java.util.Map; import java.util.Optional; import java.util.concurrent.CompletableFuture; @@ -145,6 +148,13 @@ /// ); /// ``` /// +/// ### Profiles +/// +/// A catalog key may be `name:profile`, e.g. `sift1m:label_00`, so one dataset can ship several +/// variants of its files. {@link #loadDataSet(DataSetSpec)} resolves a spec against these keys: +/// the default profile matches the bare `name` entry or `name:default`, any other profile matches +/// only `name:profile`. {@link #loadDataSet(String)} treats its argument as a literal key. +/// /// ### Metadata /// /// Dataset metadata (similarity function, load behavior) is resolved from @@ -385,8 +395,39 @@ public DataSetLoaderSimpleMFD(String catalogUrl, String localPath, boolean check } } + /// Looks up `dataSetName` as a literal catalog key, so a key written as `sift1m:label_00` is + /// found by that exact string. Metadata is looked up by the same key and then, for a + /// `name:profile` key, by the bare name. @Override public Optional loadDataSet(String dataSetName) { + int colon = dataSetName.indexOf(':'); + List metadataKeys = colon > 0 + ? List.of(dataSetName, dataSetName.substring(0, colon)) + : List.of(dataSetName); + return load(dataSetName, metadataKeys); + } + + /// Profile-aware lookup. Catalog keys are `name` or `name:profile`. A spec with the default + /// profile matches the bare `name` entry first and then `name:default`; any other profile + /// matches only `name:profile`. Metadata is looked up by the matched catalog key first and then + /// by the bare name, so datasets whose profiles share properties need a single metadata entry. + /// Nothing is downloaded or loaded unless a catalog key matches. + @Override + public Optional loadDataSet(DataSetSpec spec) { + String name = spec.getName(); + List candidates = spec.isDefaultProfile() + ? List.of(name, name + ":" + DataSetSpec.DEFAULT_PROFILE) + : List.of(name + ":" + spec.getProfile()); + for (String key : candidates) { + if (catalog.containsKey(key)) { + return load(key, key.equals(name) ? List.of(key) : List.of(key, name)); + } + } + logger.debug("No catalog entry for {} (tried {})", spec, candidates); + return Optional.empty(); + } + + private Optional load(String dataSetName, List metadataKeys) { var entry = catalog.get(dataSetName); if (entry == null) return Optional.empty(); @@ -418,19 +459,30 @@ public Optional loadDataSet(String dataSetName) { logger.info("Dataset files ready for '{}' in {}s", dataSetName, String.format("%.2f", (System.nanoTime() - startTime) / 1e9)); - var props = metadata.getProperties(dataSetName) - .orElseThrow(() -> new IllegalArgumentException( - String.format( - "Dataset '%s' was found in dataset catalog, but no metadata entry was found in dataset-metadata.yml. ", - dataSetName))); + DataSetProperties props = resolveProperties(dataSetName, metadataKeys); return Optional.of(new DataSetInfo(props, () -> { - var baseVectors = SiftLoader.readFvecs(effectiveCacheDir.resolve(baseFile).toString()); + // base vectors stay on disk behind a memory-mapped reader; DataSets' wrappers decide whether + // they are subsequently cached in heap memory + RandomAccessVectorValues baseVectors; + try { + baseVectors = new MappedFvecsRandomAccessVectorValues(effectiveCacheDir.resolve(baseFile)); + } catch (IOException e) { + throw new UncheckedIOException(e); + } var queryVectors = SiftLoader.readFvecs(effectiveCacheDir.resolve(queryFile).toString()); var gtVectors = SiftLoader.readIvecs(effectiveCacheDir.resolve(gtFile).toString()); return DataSetUtils.processDataSet(dataSetName, props, baseVectors, queryVectors, gtVectors); })); } + /// Returns the first metadata entry found under `metadataKeys`, in order, named after `dataSetName`. + private DataSetProperties resolveProperties(String dataSetName, List metadataKeys) { + return metadata.getProperties(dataSetName, metadataKeys) + .orElseThrow(() -> new IllegalArgumentException(String.format( + "Dataset '%s' was found in dataset catalog, but no metadata entry was found in dataset-metadata.yml under %s. ", + dataSetName, metadataKeys))); + } + // ======================================================================================== // CATALOG DISCOVERY & LOADING // ======================================================================================== diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetMetadataReader.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetMetadataReader.java index 3207e492f..011cfa6bf 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetMetadataReader.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetMetadataReader.java @@ -25,6 +25,7 @@ import java.nio.file.Path; import java.nio.file.Paths; import java.util.HashMap; +import java.util.List; import java.util.Map; import java.util.Optional; @@ -117,6 +118,26 @@ public Optional getProperties(String datasetKey) { }); } + /// Looks up the {@link DataSetProperties} for a dataset under several candidate keys, in order, + /// naming the result after `datasetName` rather than after the key that matched. This lets a + /// profile-specific catalog key such as `sift1m:label_00` share the metadata entry of its bare + /// name `sift1m` while still reporting the full name. + /// + /// @param datasetName the name to report for the dataset when the entry has no explicit name + /// @param lookupKeys keys to try, in order; each is resolved like {@link #getProperties(String)} + /// @return the first entry found, or empty if none of the keys match + public Optional getProperties(String datasetName, List lookupKeys) { + for (String key : lookupKeys) { + var entry = findEntry(key); + if (entry.isPresent()) { + var props = new HashMap<>(entry.get()); + props.putIfAbsent(DataSetProperties.KEY_NAME, datasetName); + return Optional.of(new DataSetProperties.PropertyMap(props)); + } + } + return Optional.empty(); + } + private Optional> findEntry(String datasetKey) { Map entry = metadata.get(datasetKey); if (entry != null) { diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpec.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpec.java new file mode 100644 index 000000000..cc64504bb --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpec.java @@ -0,0 +1,360 @@ +/* + * 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 java.util.ArrayList; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/// Identifies a dataset to load: its catalog name, an optional loader profile, and an optional +/// list of symbolic {@link DataSetWrapper}s, each with optional options, to apply after loading. +/// +/// ### Sugared form +/// A single string, parsed by {@link #parse(String)}: +/// ``` +/// cohere name only; profile "default", no wrappers +/// cohere:fast name and profile +/// cohere(mmap) name and wrappers +/// cohere:default(mmap) all three (":default" is equivalent to omitting the profile) +/// cohere(mmap, lru[grain=4096,capacityMb=512]) +/// several wrappers, applied left to right; options in brackets +/// ``` +/// The profile always precedes the parenthesised wrapper list, so `name:profile` never collides +/// with the wrapper syntax. Option keys and values may not contain `,`, `=`, `[` or `]`. +/// +/// ### Structured form +/// A YAML map, accepted by {@link #from(Object)} wherever a dataset list entry may appear. Each +/// wrapper is a plain name, a map keyed by the wrapper name whose value holds its options, or a +/// map with a `name` key and the options alongside: +/// ```yaml +/// - name: cohere +/// profile: default # optional +/// wrappers: # optional +/// - mmap +/// - lru: { grain: 4096, capacityMb: 512 } +/// - name: lru +/// grain: 4096 +/// ``` +/// +/// A missing profile is always reported as {@value #DEFAULT_PROFILE}. Loaders that do not +/// understand profiles accept the default profile and reject any other; see +/// {@link DataSetLoader#loadDataSet(DataSetSpec)}. {@link #toString()} renders the canonical +/// sugared form, which round-trips through {@link #parse(String)} including wrapper options. +public final class DataSetSpec { + /// The profile assumed when none is given. + public static final String DEFAULT_PROFILE = "default"; + + private static final Pattern SUGAR = Pattern.compile( + "^\\s*([^:()\\s\\[\\]][^:()\\[\\]]*?)\\s*(?::\\s*([^:()\\[\\]]+?)\\s*)?(?:\\(([^()]*)\\)\\s*)?$"); + private static final Pattern WRAPPER_TOKEN = Pattern.compile( + "^\\s*([^\\[\\],=\\s]+)\\s*(?:\\[([^\\[\\]]*)\\]\\s*)?$"); + + /// One wrapper to apply: its registered name and its options as written. + public static final class WrapperSpec { + private final String name; + private final Map options; + + /// @param name the registered wrapper name; must not be blank + /// @param options option key/value pairs, or null for none; values are kept as strings + public WrapperSpec(String name, Map options) { + if (name == null || name.isBlank()) { + throw new IllegalArgumentException("Wrapper name must not be blank"); + } + this.name = name.trim(); + Map copy = new LinkedHashMap<>(); + if (options != null) { + for (var e : options.entrySet()) { + String key = e.getKey() == null ? "" : e.getKey().trim(); + if (key.isEmpty()) { + throw new IllegalArgumentException("Wrapper '" + this.name + "' has an option with an empty key"); + } + String value = e.getValue() == null ? "" : String.valueOf(e.getValue()).trim(); + for (String forbidden : new String[] {",", "=", "[", "]"}) { + if (key.contains(forbidden) || value.contains(forbidden)) { + throw new IllegalArgumentException("Wrapper '" + this.name + "' option '" + key + "=" + value + + "' may not contain '" + forbidden + "'"); + } + } + copy.put(key, value); + } + } + this.options = Collections.unmodifiableMap(copy); + } + + /// Parses a token such as `lru` or `lru[grain=4096,capacityMb=512]`. + static WrapperSpec parse(String token) { + Matcher m = WRAPPER_TOKEN.matcher(token); + if (!m.matches()) { + throw new IllegalArgumentException("Malformed wrapper '" + token.trim() + "'; expected name or name[key=value,...]"); + } + Map options = new LinkedHashMap<>(); + if (m.group(2) != null && !m.group(2).isBlank()) { + for (String pair : m.group(2).split(",")) { + int eq = pair.indexOf('='); + if (eq < 0) { + throw new IllegalArgumentException("Malformed wrapper option '" + pair.trim() + "' in '" + token.trim() + "'; expected key=value"); + } + options.put(pair.substring(0, eq).trim(), pair.substring(eq + 1).trim()); + } + } + return new WrapperSpec(m.group(1), options); + } + + /// @return the registered wrapper name + public String getName() { + return name; + } + + /// @return the options as written, in order; empty when none were given + public Map getOptions() { + return options; + } + + /// @return `name` or `name[key=value,...]` + @Override + public String toString() { + if (options.isEmpty()) { + return name; + } + StringBuilder sb = new StringBuilder(name).append('['); + boolean first = true; + for (var e : options.entrySet()) { + if (!first) sb.append(','); + first = false; + sb.append(e.getKey()).append('=').append(e.getValue()); + } + return sb.append(']').toString(); + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (!(o instanceof WrapperSpec)) return false; + WrapperSpec that = (WrapperSpec) o; + return name.equals(that.name) && options.equals(that.options); + } + + @Override + public int hashCode() { + return Objects.hash(name, options); + } + } + + private final String name; + private final String profile; + private final List wrappers; + + /// Creates a spec. Blank profile means {@value #DEFAULT_PROFILE}. + /// + /// @param name the dataset name; must not be blank + /// @param profile the loader profile, or null for the default + /// @param wrappers wrappers in application order, or null for none + public DataSetSpec(String name, String profile, List wrappers) { + if (name == null || name.isBlank()) { + throw new IllegalArgumentException("Dataset name must not be blank"); + } + this.name = name.trim(); + this.profile = (profile == null || profile.isBlank()) ? DEFAULT_PROFILE : profile.trim(); + this.wrappers = wrappers == null ? List.of() : List.copyOf(wrappers); + } + + /// Parses the sugared string form described in the class documentation. + /// + /// @param sugared e.g. `cohere`, `cohere:fast`, `cohere(mmap)`, or `cohere:fast(mmap, lru[grain=4096])` + /// @return the parsed spec + /// @throws IllegalArgumentException if the string is not in the sugared form + public static DataSetSpec parse(String sugared) { + if (sugared == null) { + throw new IllegalArgumentException("Dataset spec must not be null"); + } + Matcher m = SUGAR.matcher(sugared); + if (!m.matches()) { + throw new IllegalArgumentException("Malformed dataset spec '" + sugared + + "'; expected name, name:profile, name(wrapper,...), or name:profile(wrapper,...)"); + } + return new DataSetSpec(m.group(1), m.group(2), parseWrapperList(m.group(3))); + } + + /// Splits a wrapper list on the commas that are not inside `[...]` and parses each token. + static List parseWrapperList(String list) { + List wrappers = new ArrayList<>(); + if (list == null) { + return wrappers; + } + int depth = 0; + int start = 0; + for (int i = 0; i <= list.length(); i++) { + char c = i < list.length() ? list.charAt(i) : ','; + if (c == '[') { + depth++; + } else if (c == ']') { + depth--; + if (depth < 0) { + throw new IllegalArgumentException("Unbalanced ']' in wrapper list '" + list + "'"); + } + } else if (c == ',' && depth == 0) { + String token = list.substring(start, i); + if (!token.isBlank()) { + wrappers.add(WrapperSpec.parse(token)); + } + start = i + 1; + } + } + if (depth != 0) { + throw new IllegalArgumentException("Unbalanced '[' in wrapper list '" + list + "'"); + } + return wrappers; + } + + /// Converts a dataset list entry from YAML: a string in sugared form, or a map with `name` + /// and optional `profile` and `wrappers` keys. `wrappers` may be a list or a single string; + /// each list element is a wrapper name (optionally with bracketed options), a map keyed by the + /// wrapper name whose value is its option map, or a map with a `name` key and options alongside. + /// + /// @param item the YAML value + /// @return the spec + /// @throws IllegalArgumentException if the value is neither form, or the map lacks a name + @SuppressWarnings("unchecked") + public static DataSetSpec from(Object item) { + if (item instanceof DataSetSpec) { + return (DataSetSpec) item; + } + if (item instanceof String) { + return parse((String) item); + } + if (item instanceof Map) { + Map map = (Map) item; + Object name = map.get("name"); + if (name == null) { + throw new IllegalArgumentException("Structured dataset entry is missing 'name': " + map); + } + for (String key : map.keySet()) { + if (!key.equals("name") && !key.equals("profile") && !key.equals("wrappers")) { + throw new IllegalArgumentException("Unknown key '" + key + "' in dataset entry " + map + + "; known keys: name, profile, wrappers"); + } + } + Object profile = map.get("profile"); + Object wrappers = map.get("wrappers"); + List wrapperSpecs; + if (wrappers == null) { + wrapperSpecs = List.of(); + } else if (wrappers instanceof String) { + wrapperSpecs = parseWrapperList((String) wrappers); + } else if (wrappers instanceof List) { + wrapperSpecs = new ArrayList<>(); + for (Object w : (List) wrappers) { + wrapperSpecs.add(wrapperFrom(w)); + } + } else { + throw new IllegalArgumentException("'wrappers' must be a list or string in dataset entry " + map); + } + return new DataSetSpec(name.toString(), profile == null ? null : profile.toString(), wrapperSpecs); + } + throw new IllegalArgumentException("Dataset entry must be a string or a map, got: " + item); + } + + @SuppressWarnings("unchecked") + private static WrapperSpec wrapperFrom(Object item) { + if (item instanceof String) { + return WrapperSpec.parse((String) item); + } + if (item instanceof Map) { + Map map = (Map) item; + Object name = map.get("name"); + if (name != null) { + Map options = new LinkedHashMap<>(map); + options.remove("name"); + return new WrapperSpec(name.toString(), options); + } + if (map.size() == 1) { + var entry = map.entrySet().iterator().next(); + Object value = entry.getValue(); + if (value == null) { + return new WrapperSpec(entry.getKey(), null); + } + if (value instanceof Map) { + return new WrapperSpec(entry.getKey(), (Map) value); + } + throw new IllegalArgumentException("Options for wrapper '" + entry.getKey() + "' must be a map, got: " + value); + } + throw new IllegalArgumentException("Wrapper entry must be a name, {name: options-map}, or {name: ..., option: value}; got: " + map); + } + throw new IllegalArgumentException("Wrapper entry must be a string or a map, got: " + item); + } + + /// @return the dataset name as known to the loaders + public String getName() { + return name; + } + + /// @return the loader profile; {@value #DEFAULT_PROFILE} when none was given + public String getProfile() { + return profile; + } + + /// @return true iff the profile is {@value #DEFAULT_PROFILE} + public boolean isDefaultProfile() { + return DEFAULT_PROFILE.equals(profile); + } + + /// @return wrappers in application order, with their options; empty when none were given + public List getWrappers() { + return wrappers; + } + + /// @return true iff at least one wrapper was given + public boolean hasWrappers() { + return !wrappers.isEmpty(); + } + + /// @return the canonical sugared form: the name, `:profile` unless default, and `(w1,w2[k=v])` if any wrappers + @Override + public String toString() { + StringBuilder sb = new StringBuilder(name); + if (!isDefaultProfile()) { + sb.append(':').append(profile); + } + if (hasWrappers()) { + sb.append('('); + for (int i = 0; i < wrappers.size(); i++) { + if (i > 0) sb.append(','); + sb.append(wrappers.get(i)); + } + sb.append(')'); + } + return sb.toString(); + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (!(o instanceof DataSetSpec)) return false; + DataSetSpec that = (DataSetSpec) o; + return name.equals(that.name) && profile.equals(that.profile) && wrappers.equals(that.wrappers); + } + + @Override + public int hashCode() { + return Objects.hash(name, profile, wrappers); + } +} 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..67c117777 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 @@ -16,6 +16,7 @@ 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.VectorUtil; import io.github.jbellis.jvector.vector.types.VectorFloat; @@ -42,10 +43,22 @@ public class DataSetUtils { /** * Processes a dataset using the configured load behavior from the dataset metadata. + *

+ * With {@link DataSetProperties.LoadBehavior#NO_SCRUB} the base vectors are used exactly as given, so a + * memory-mapped reader stays memory-mapped. {@link DataSetProperties.LoadBehavior#LEGACY_SCRUB} must + * inspect and rewrite every vector, so it first copies them into heap memory via + * {@link InMemoryCachedDataSet#readAllVectors(RandomAccessVectorValues)}. + * + * @param pathStr the dataset name + * @param props the dataset properties, supplying the similarity function and load behavior + * @param baseVectors the base vectors, from any kind of reader + * @param queryVectors the query vectors + * @param groundTruth one neighbor list per query vector + * @return the dataset */ public static DataSet processDataSet(String pathStr, DataSetProperties props, - List> baseVectors, + RandomAccessVectorValues baseVectors, List> queryVectors, List> groundTruth) { var vsf = props.similarityFunction() @@ -56,7 +69,7 @@ public static DataSet processDataSet(String pathStr, case NO_SCRUB: return new SimpleDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); case LEGACY_SCRUB: - return legacyScrubDataSet(pathStr, vsf, baseVectors, queryVectors, groundTruth); + return legacyScrubDataSet(pathStr, vsf, InMemoryCachedDataSet.readAllVectors(baseVectors), queryVectors, groundTruth); default: throw new IllegalArgumentException("Unsupported load behavior: " + props.loadBehavior()); } @@ -64,7 +77,7 @@ public static DataSet processDataSet(String pathStr, /** * @deprecated Benchmark loaders should use - * {@link #processDataSet(String, DataSetProperties, List, List, List)} + * {@link #processDataSet(String, DataSetProperties, RandomAccessVectorValues, List, List)} * so that load behavior is controlled explicitly by dataset metadata. */ @Deprecated(forRemoval = true) diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapper.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapper.java new file mode 100644 index 000000000..3f52ae9c0 --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapper.java @@ -0,0 +1,109 @@ +/* + * 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.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorSimilarityFunction; +import io.github.jbellis.jvector.vector.types.VectorFloat; + +import java.util.List; +import java.util.Map; + +/// A {@link DataSet} layered over an origin dataset. +/// +/// Wrappers change *how* a dataset's vectors are held or served without changing *what* they +/// are: every accessor delegates to {@link #getOrigin()} unless the wrapper overrides it. The +/// built-in wrappers re-home the base vectors: fully in heap memory ({@link InMemoryCachedDataSet}), +/// in a memory-mapped file ({@link MMapCachedDataSet}), or behind a bounded grain-resolution LRU +/// cache for datasets larger than memory ({@link LruCachedDataSet}). +/// +/// Wrappers are applied by {@link DataSets} through a {@link Provider}, either passed explicitly +/// or resolved by name from {@link DataSets#wrapperProviders} when a dataset is named with a +/// symbolic wrapper list such as `cohere-english-v3-100k(mmap)` (see {@link DataSetSpec}). +/// +/// @see DataSets +/// @see DataSetSpec +public interface DataSetWrapper extends DataSet { + + /// @return the dataset this wrapper is layered over + DataSet getOrigin(); + + @Override + default int getDimension() { + return getOrigin().getDimension(); + } + + @Override + default RandomAccessVectorValues getBaseRavv() { + return getOrigin().getBaseRavv(); + } + + @Override + default String getName() { + return getOrigin().getName(); + } + + @Override + default VectorSimilarityFunction getSimilarityFunction() { + return getOrigin().getSimilarityFunction(); + } + + @Override + default List> getQueryVectors() { + return getOrigin().getQueryVectors(); + } + + @Override + default List> getGroundTruth() { + return getOrigin().getGroundTruth(); + } + + /// Creates a {@link DataSetWrapper} over an origin dataset. + /// + /// Providers should be idempotent for their own wrapper type: wrapping an already-wrapped + /// dataset of the same kind returns it unchanged rather than stacking a redundant layer. + @FunctionalInterface + interface Provider { + /// @param origin the dataset to wrap + /// @return the wrapped dataset + DataSetWrapper wrap(DataSet origin); + } + + /// Builds a {@link Provider} from the options written on a {@link DataSetSpec.WrapperSpec}. + /// This is what {@link DataSets#wrapperProviders} registers under each symbolic wrapper name. + @FunctionalInterface + interface Factory { + /// @param options the wrapper's options as written, possibly empty + /// @return a provider configured with those options + /// @throws IllegalArgumentException if an option is unknown or malformed + Provider provider(Map options); + + /// Wraps a provider that takes no options; any option given is rejected by name. + /// + /// @param wrapperName the symbolic name, for the error message + /// @param provider the provider to return + /// @return a factory that yields `provider` when given no options + static Factory optionless(String wrapperName, Provider provider) { + return options -> { + if (!options.isEmpty()) { + throw new IllegalArgumentException("Dataset wrapper '" + wrapperName + "' takes no options, got " + options.keySet()); + } + return provider; + }; + } + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSets.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSets.java index 94b2e1c77..a51d83ed6 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSets.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSets.java @@ -22,17 +22,32 @@ import java.security.InvalidParameterException; import java.util.ArrayList; import java.util.Collection; +import java.util.LinkedHashMap; import java.util.List; +import java.util.Map; import java.util.Optional; -/// Facade for locating datasets across multiple {@link DataSetLoader} implementations. +/// Facade for locating datasets across multiple {@link DataSetLoader} implementations and +/// layering {@link DataSetWrapper}s over what they load. /// /// Returns a {@link DataSetInfo} handle whose vector data is loaded lazily on the first /// call to {@link DataSetInfo#getDataSet()}, allowing callers to inspect dataset metadata /// (name, similarity function) without incurring the cost of reading vectors into memory. /// +/// ### Wrappers +/// Every dataset name accepted here may be a {@link DataSetSpec} in sugared form, e.g. +/// `cohere-english-v3-100k(mmap)` or `cohere-english-v3-100k(mmap,lru[grain=4096,capacityMb=512])`. +/// Symbolic wrappers are resolved through {@link #wrapperProviders}, which hands each wrapper its +/// bracketed (or, in YAML, mapped) options. When a spec names no wrappers, +/// {@link #defaultWrappers} apply, which cache base vectors in heap memory so that datasets +/// behave as they did when loaders read everything into lists. Datasets larger than memory name +/// `mmap` (serve straight from the mapped file) or `lru` (a bounded grain cache in front of it) +/// instead, which replaces that default. Passing wrapper providers +/// explicitly replaces both; a spec that also names wrappers is rejected in that case. +/// /// @see DataSetInfo /// @see DataSetLoader +/// @see DataSetSpec public class DataSets { private static final Logger logger = LoggerFactory.getLogger(DataSets.class); @@ -48,34 +63,131 @@ public class DataSets { }}; - /// Loads a dataset by name using the {@link #defaultLoaders}. + /// Symbolic wrapper names, as used in `name(wrapper,...)` and structured `wrappers:` lists, + /// mapped to the factory that builds each from its options. Additional wrappers may be + /// registered here; a wrapper without options registers via {@link DataSetWrapper.Factory#optionless}. + public static final Map wrapperProviders = new LinkedHashMap<>() {{ + put(InMemoryCachedDataSet.WRAPPER_NAME, InMemoryCachedDataSet.FACTORY); + put(MMapCachedDataSet.WRAPPER_NAME, MMapCachedDataSet.FACTORY); + put(LruCachedDataSet.WRAPPER_NAME, LruCachedDataSet.FACTORY); + }}; + + /// Wrappers applied when a dataset spec names none: the base vectors are cached in heap memory. + public static final List defaultWrappers = new ArrayList<>(List.of(InMemoryCachedDataSet.PROVIDER)); + + /// Loads a dataset by name or sugared spec using the {@link #defaultLoaders} and either the + /// spec's wrappers or the {@link #defaultWrappers}. /// - /// @param dataSetName the logical dataset name (e.g. {@code "ada002-100k"}) + /// @param dataSetName the logical dataset name (e.g. {@code "ada002-100k"}), optionally with profile and wrappers /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name public static Optional loadDataSet(String dataSetName) { - return loadDataSet(dataSetName, defaultLoaders); + return loadDataSet(DataSetSpec.parse(dataSetName)); + } + + /// Loads a dataset by spec using the {@link #defaultLoaders} and either the spec's wrappers or + /// the {@link #defaultWrappers}. + /// + /// @param spec the dataset name, profile, and wrapper names + /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name + public static Optional loadDataSet(DataSetSpec spec) { + return loadDataSet(spec, defaultLoaders); } - /// Loads a dataset by name, trying each loader in order until one matches. + /// Loads a dataset by name or sugared spec, trying each loader in order until one matches, and + /// applying either the spec's wrappers or the {@link #defaultWrappers}. /// - /// @param dataSetName the logical dataset name (e.g. {@code "ada002-100k"}) + /// @param dataSetName the logical dataset name (e.g. {@code "ada002-100k"}), optionally with profile and wrappers /// @param loaders the loaders to try, in priority order /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name public static Optional loadDataSet(String dataSetName, Collection loaders) { - logger.info("loading dataset [{}]", dataSetName); - if (dataSetName.endsWith(".hdf5")) { - throw new InvalidParameterException("DataSet names are not meant to be file names. Did you mean " + dataSetName.replace(".hdf5", "") + "? "); + return loadDataSet(DataSetSpec.parse(dataSetName), loaders); + } + + /// Loads a dataset by spec, trying each loader in order until one matches, and applying either + /// the spec's wrappers (resolved through {@link #wrapperProviders}) or the {@link #defaultWrappers}. + /// + /// @param spec the dataset name, profile, and wrapper names + /// @param loaders the loaders to try, in priority order + /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name + /// @throws IllegalArgumentException if the spec names a wrapper that is not registered + public static Optional loadDataSet(DataSetSpec spec, Collection loaders) { + List wrappers = spec.hasWrappers() ? resolveWrappers(spec.getWrappers()) : defaultWrappers; + return loadDataSet(new DataSetSpec(spec.getName(), spec.getProfile(), null), loaders, wrappers); + } + + /// Loads a dataset by name, trying each loader in order until one matches, then applies exactly + /// the given wrapper providers, in order, when the dataset is first materialised. + /// + /// @param dataSetName the logical dataset name, optionally with a profile; it must not name wrappers + /// @param loaders the loaders to try, in priority order + /// @param wrappers the wrappers to layer over the loaded dataset, outermost last; may be empty + /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name + public static Optional loadDataSet(String dataSetName, + Collection loaders, + Collection wrappers) { + return loadDataSet(DataSetSpec.parse(dataSetName), loaders, wrappers); + } + + /// Loads a dataset by spec, trying each loader in order until one matches, then applies exactly + /// the given wrapper providers, in order, when the dataset is first materialised. + /// + /// @param spec the dataset name and profile; it must not name wrappers, since `wrappers` replaces them + /// @param loaders the loaders to try, in priority order + /// @param wrappers the wrappers to layer over the loaded dataset, outermost last; may be empty + /// @return a lazy {@link DataSetInfo} handle, or empty if no loader recognises the name + /// @throws IllegalArgumentException if the spec names wrappers as well + public static Optional loadDataSet(DataSetSpec spec, + Collection loaders, + Collection wrappers) { + if (spec.hasWrappers()) { + throw new IllegalArgumentException("Dataset spec '" + spec + "' names wrappers, but wrapper providers were also given explicitly; use one or the other"); + } + logger.info("loading dataset [{}]", spec); + if (spec.getName().endsWith(".hdf5")) { + throw new InvalidParameterException("DataSet names are not meant to be file names. Did you mean " + spec.getName().replace(".hdf5", "") + "? "); } for (DataSetLoader loader : loaders) { logger.trace("trying loader [{}]", loader.getClass().getSimpleName()); - Optional dataSetLoaded = loader.loadDataSet(dataSetName); + Optional dataSetLoaded = loader.loadDataSet(spec); if (dataSetLoaded.isPresent()) { - logger.info("dataset [{}] found with loader [{}]", dataSetName, loader.getClass().getSimpleName()); - return dataSetLoaded; + logger.info("dataset [{}] found with loader [{}]", spec, loader.getClass().getSimpleName()); + return Optional.of(wrap(dataSetLoaded.get(), wrappers)); } } - logger.warn("Unable to find dataset [{}] with any dataset loader.", dataSetName); + logger.warn("Unable to find dataset [{}] with any dataset loader.", spec); return Optional.empty(); } + + /// Resolves symbolic wrappers through {@link #wrapperProviders}, preserving order and handing + /// each wrapper's options to its factory. + /// + /// @param wrappers wrappers as written in a spec + /// @return the configured providers + /// @throws IllegalArgumentException if a name is not registered or a factory rejects its options + public static List resolveWrappers(List wrappers) { + List providers = new ArrayList<>(wrappers.size()); + for (DataSetSpec.WrapperSpec wrapper : wrappers) { + DataSetWrapper.Factory factory = wrapperProviders.get(wrapper.getName()); + if (factory == null) { + throw new IllegalArgumentException("Unknown dataset wrapper '" + wrapper.getName() + "'; known wrappers: " + wrapperProviders.keySet()); + } + providers.add(factory.provider(wrapper.getOptions())); + } + return providers; + } + + private static DataSetInfo wrap(DataSetInfo info, Collection wrappers) { + if (wrappers.isEmpty()) { + return info; + } + List providers = List.copyOf(wrappers); + return new DataSetInfo(info, () -> { + DataSet ds = info.getDataSet(); + for (DataSetWrapper.Provider provider : providers) { + ds = provider.wrap(ds); + } + return ds; + }); + } } diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/InMemoryCachedDataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/InMemoryCachedDataSet.java new file mode 100644 index 000000000..31caa50e8 --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/InMemoryCachedDataSet.java @@ -0,0 +1,131 @@ +/* + * 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.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import java.util.Arrays; +import java.util.List; +import java.util.stream.IntStream; + +/// A {@link DataSetWrapper} whose base vectors are fully cached in heap memory. +/// +/// On construction the origin's base {@link RandomAccessVectorValues} is read in its entirety +/// through {@link RandomAccessVectorValues#range(int, int)} views, one chunk per task, in +/// parallel on the common pool. Each vector is copied into a fresh heap vector, so the cache is +/// independent of the origin's storage (a memory-mapped file, for instance) and serves reads at +/// plain list-lookup cost. +/// +/// An origin whose base RAVV is already a {@link ListRandomAccessVectorValues} is adopted as-is: +/// it is heap-resident already, and copying it would only double the footprint. +/// +/// Registered in {@link DataSets#wrapperProviders} under the name {@value #WRAPPER_NAME}. +public final class InMemoryCachedDataSet implements DataSetWrapper { + private static final Logger logger = LoggerFactory.getLogger(InMemoryCachedDataSet.class); + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + /// Symbolic wrapper name, as used in `dataset(memory)` and `wrappers: [memory]`. + public static final String WRAPPER_NAME = "memory"; + + /// Provider that caches a dataset's base vectors in memory; idempotent for already-cached datasets. + public static final DataSetWrapper.Provider PROVIDER = InMemoryCachedDataSet::of; + + /// Registry entry for {@value #WRAPPER_NAME}; this wrapper takes no options. + public static final DataSetWrapper.Factory FACTORY = DataSetWrapper.Factory.optionless(WRAPPER_NAME, PROVIDER); + + /// Vectors copied per parallel task. + static final int DEFAULT_CHUNK_SIZE = 8 * 1024; + + private final DataSet origin; + private final RandomAccessVectorValues baseRavv; + + /// Returns `origin` itself if it is already an {@link InMemoryCachedDataSet}, otherwise a new cache over it. + /// + /// @param origin the dataset whose base vectors should be heap-resident + /// @return an in-memory view of `origin` + public static DataSetWrapper of(DataSet origin) { + if (origin instanceof InMemoryCachedDataSet) { + return (InMemoryCachedDataSet) origin; + } + return new InMemoryCachedDataSet(origin); + } + + private InMemoryCachedDataSet(DataSet origin) { + this.origin = origin; + RandomAccessVectorValues source = origin.getBaseRavv(); + if (source instanceof ListRandomAccessVectorValues) { + this.baseRavv = source; + logger.info("Base vectors of '{}' are already heap-resident; adopting {} vectors as-is", origin.getName(), source.size()); + return; + } + long start = System.nanoTime(); + List> vectors = readAllVectors(source); + this.baseRavv = new ListRandomAccessVectorValues(vectors, source.dimension()); + double mb = (double) vectors.size() * source.dimension() * Float.BYTES / (1024.0 * 1024.0); + logger.info("Cached {} base vectors ({} MB) of '{}' in memory in {}s", + vectors.size(), String.format("%.1f", mb), origin.getName(), + String.format("%.2f", (System.nanoTime() - start) / 1e9)); + } + + /// Copies every vector of `source` into a new heap-resident list, reading through ranged views + /// in parallel chunks of {@value #DEFAULT_CHUNK_SIZE} vectors. + /// + /// @param source the vectors to copy; may be value-shared, since each chunk reads through its own {@link RandomAccessVectorValues#copy()} + /// @return a fixed-size list with one independent vector per ordinal of `source` + public static List> readAllVectors(RandomAccessVectorValues source) { + return readAllVectors(source, DEFAULT_CHUNK_SIZE); + } + + static List> readAllVectors(RandomAccessVectorValues source, int chunkSize) { + int count = source.size(); + int dimension = source.dimension(); + VectorFloat[] out = new VectorFloat[count]; + int chunks = (count + chunkSize - 1) / chunkSize; + IntStream.range(0, chunks).parallel().forEach(chunk -> { + int start = chunk * chunkSize; + int end = Math.min(count, start + chunkSize); + RandomAccessVectorValues slice = source.range(start, end).copy(); + for (int i = 0; i < slice.size(); i++) { + VectorFloat v = vts.createFloatVector(dimension); + slice.getVectorInto(i, v, 0); + out[start + i] = v; + } + }); + return Arrays.asList(out); + } + + @Override + public DataSet getOrigin() { + return origin; + } + + @Override + public RandomAccessVectorValues getBaseRavv() { + return baseRavv; + } + + @Override + public int getDimension() { + return baseRavv.dimension(); + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/LruCachedDataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/LruCachedDataSet.java new file mode 100644 index 000000000..f4115052e --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/LruCachedDataSet.java @@ -0,0 +1,172 @@ +/* + * 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.example.util.LruGrainRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import java.util.Map; + +/// A {@link DataSetWrapper} that serves base vectors through a bounded, grain-resolution LRU cache +/// ({@link LruGrainRandomAccessVectorValues}) in front of the origin's reader. +/// +/// This is the wrapper for datasets larger than memory: the origin (typically the loader's +/// memory-mapped fvecs reader) stays on disk, and only the most recently used grains of +/// `grainSize` consecutive vectors are held in heap memory, up to `capacityBytes` in total. +/// +/// Registered in {@link DataSets#wrapperProviders} under the name {@value #WRAPPER_NAME}, with the +/// options {@value #GRAIN_OPTION} and {@value #CAPACITY_MB_OPTION}. An option that is not given uses +/// a default, each overridable with a system property: +/// +/// - `{@value #GRAIN_PROPERTY}` — vectors per grain, default {@value #DEFAULT_GRAIN_SIZE} +/// - `{@value #CAPACITY_MB_PROPERTY}` — cache capacity in MiB, default one quarter of the maximum heap +/// +/// Programmatic users pick their own numbers with {@link #provider(int, long)}. +public final class LruCachedDataSet implements DataSetWrapper { + private static final Logger logger = LoggerFactory.getLogger(LruCachedDataSet.class); + + /// Symbolic wrapper name, as used in `dataset(lru)` and `wrappers: [lru]`. + public static final String WRAPPER_NAME = "lru"; + + /// System property naming the vectors per grain for the symbolic wrapper. + public static final String GRAIN_PROPERTY = "jvector.dataset.lru.grain"; + + /// System property naming the cache capacity in MiB for the symbolic wrapper. + public static final String CAPACITY_MB_PROPERTY = "jvector.dataset.lru.capacityMb"; + + /// Vectors per grain when {@value #GRAIN_PROPERTY} is not set. + public static final int DEFAULT_GRAIN_SIZE = 1024; + + /// Provider using the property-driven defaults; idempotent for already-wrapped datasets. + public static final DataSetWrapper.Provider PROVIDER = LruCachedDataSet::of; + + /// Option key for vectors per grain, e.g. `lru[grain=4096]` or `lru: { grain: 4096 }`. + public static final String GRAIN_OPTION = "grain"; + + /// Option key for cache capacity in MiB, e.g. `lru[capacityMb=512]` or `lru: { capacityMb: 512 }`. + public static final String CAPACITY_MB_OPTION = "capacityMb"; + + /// Registry entry for {@value #WRAPPER_NAME}: accepts {@value #GRAIN_OPTION} and + /// {@value #CAPACITY_MB_OPTION}, each falling back to the property-driven default when absent. + public static final DataSetWrapper.Factory FACTORY = LruCachedDataSet::provider; + + /// Builds a provider from wrapper options. + /// + /// @param options {@value #GRAIN_OPTION} (positive integer) and/or {@value #CAPACITY_MB_OPTION} (positive integer) + /// @return a provider with those settings, defaults filling any absent option + /// @throws IllegalArgumentException on an unknown key or a non-positive or non-numeric value + public static DataSetWrapper.Provider provider(Map options) { + int grainSize = defaultGrainSize(); + long capacityBytes = defaultCapacityBytes(); + for (var e : options.entrySet()) { + switch (e.getKey()) { + case GRAIN_OPTION: + grainSize = (int) positive(e.getKey(), e.getValue(), Integer.MAX_VALUE); + break; + case CAPACITY_MB_OPTION: + capacityBytes = positive(e.getKey(), e.getValue(), Long.MAX_VALUE / (1024 * 1024)) * 1024 * 1024; + break; + default: + throw new IllegalArgumentException("Unknown option '" + e.getKey() + "' for dataset wrapper '" + WRAPPER_NAME + + "'; known options: " + GRAIN_OPTION + ", " + CAPACITY_MB_OPTION); + } + } + if (options.isEmpty()) { + return PROVIDER; + } + return provider(grainSize, capacityBytes); + } + + private static long positive(String key, String value, long max) { + long parsed; + try { + parsed = Long.parseLong(value); + } catch (NumberFormatException e) { + throw new IllegalArgumentException("Option '" + key + "' of dataset wrapper '" + WRAPPER_NAME + "' must be an integer, got '" + value + "'"); + } + if (parsed <= 0 || parsed > max) { + throw new IllegalArgumentException("Option '" + key + "' of dataset wrapper '" + WRAPPER_NAME + "' must be between 1 and " + max + ", got " + parsed); + } + return parsed; + } + + private final DataSet origin; + private final LruGrainRandomAccessVectorValues baseRavv; + + /// Returns `origin` itself if it is already an {@link LruCachedDataSet}, otherwise a new cache over it + /// sized by {@link #defaultGrainSize()} and {@link #defaultCapacityBytes()}. + /// + /// @param origin the dataset whose base vectors should be served through a bounded cache + /// @return an LRU-cached view of `origin` + public static DataSetWrapper of(DataSet origin) { + if (origin instanceof LruCachedDataSet) { + return (LruCachedDataSet) origin; + } + return new LruCachedDataSet(origin, defaultGrainSize(), defaultCapacityBytes()); + } + + /// @param grainSize vectors per grain + /// @param capacityBytes heap bytes of vector data to keep resident + /// @return a provider building caches with exactly these parameters + public static DataSetWrapper.Provider provider(int grainSize, long capacityBytes) { + return origin -> new LruCachedDataSet(origin, grainSize, capacityBytes); + } + + /// @return `{@value #GRAIN_PROPERTY}` or {@value #DEFAULT_GRAIN_SIZE} + public static int defaultGrainSize() { + return Integer.getInteger(GRAIN_PROPERTY, DEFAULT_GRAIN_SIZE); + } + + /// @return `{@value #CAPACITY_MB_PROPERTY}` in bytes, or one quarter of the maximum heap + public static long defaultCapacityBytes() { + Long mb = Long.getLong(CAPACITY_MB_PROPERTY); + return mb != null ? mb * 1024 * 1024 : Runtime.getRuntime().maxMemory() / 4; + } + + /// Creates a cache of `capacityBytes` worth of `grainSize`-vector grains over `origin`. + /// + /// @param origin the dataset to wrap + /// @param grainSize vectors per grain; must be positive + /// @param capacityBytes heap bytes of vector data to keep resident; at least one grain is always kept + public LruCachedDataSet(DataSet origin, int grainSize, long capacityBytes) { + this.origin = origin; + RandomAccessVectorValues source = origin.getBaseRavv(); + long grainBytes = (long) grainSize * source.dimension() * Float.BYTES; + int maxGrains = (int) Math.max(1, Math.min(Integer.MAX_VALUE, capacityBytes / grainBytes)); + this.baseRavv = new LruGrainRandomAccessVectorValues(source, grainSize, maxGrains); + logger.info("LRU cache over '{}': {} vectors in grains of {} ({} MB each), keeping at most {} of {} grains", + origin.getName(), source.size(), grainSize, String.format("%.1f", grainBytes / (1024.0 * 1024.0)), + maxGrains, baseRavv.grainCount()); + } + + @Override + public DataSet getOrigin() { + return origin; + } + + @Override + public LruGrainRandomAccessVectorValues getBaseRavv() { + return baseRavv; + } + + @Override + public int getDimension() { + return baseRavv.dimension(); + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/MMapCachedDataSet.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/MMapCachedDataSet.java new file mode 100644 index 000000000..2ac2b298e --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/benchmarks/datasets/MMapCachedDataSet.java @@ -0,0 +1,150 @@ +/* + * 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.example.util.MappedFvecsRandomAccessVectorValues; +import io.github.jbellis.jvector.example.util.SiftLoader; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import java.io.IOException; +import java.io.UncheckedIOException; +import java.nio.file.AtomicMoveNotSupportedException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.nio.file.StandardCopyOption; + +/// A {@link DataSetWrapper} whose base vectors are served from a memory-mapped fvecs file. +/// +/// If the origin's base vectors are already a {@link MappedFvecsRandomAccessVectorValues} (as +/// produced by {@link DataSetLoaderSimpleMFD}), they are adopted as-is. Otherwise the base +/// vectors are spilled to `/-x.fvecs` and mapped from there, so +/// the heap no longer holds them. An existing spill file of exactly the expected size is reused +/// rather than rewritten, both because a mapped file cannot be modified on Windows and because +/// the spill is then paid once per dataset rather than once per run; delete the file to force a +/// fresh spill. +/// +/// The default cache directory is `$DATASET_CACHE_DIR/mmap` when that variable is set, else +/// `dataset_cache/mmap` relative to the working directory. +/// +/// Registered in {@link DataSets#wrapperProviders} under the name {@value #WRAPPER_NAME}. +public final class MMapCachedDataSet implements DataSetWrapper { + private static final Logger logger = LoggerFactory.getLogger(MMapCachedDataSet.class); + private static final String ENV_DATASET_CACHE_DIR = "DATASET_CACHE_DIR"; + + /// Symbolic wrapper name, as used in `dataset(mmap)` and `wrappers: [mmap]`. + public static final String WRAPPER_NAME = "mmap"; + + /// Provider that memory-maps a dataset's base vectors under {@link #defaultCacheDir()}; idempotent for already-mapped datasets. + public static final DataSetWrapper.Provider PROVIDER = MMapCachedDataSet::of; + + /// Registry entry for {@value #WRAPPER_NAME}; this wrapper takes no options. + public static final DataSetWrapper.Factory FACTORY = DataSetWrapper.Factory.optionless(WRAPPER_NAME, PROVIDER); + + private final DataSet origin; + private final RandomAccessVectorValues baseRavv; + + /// Returns `origin` itself if it is already an {@link MMapCachedDataSet}, otherwise a new mapped view + /// using {@link #defaultCacheDir()} for any spill file. + /// + /// @param origin the dataset whose base vectors should be memory-mapped + /// @return a memory-mapped view of `origin` + public static DataSetWrapper of(DataSet origin) { + if (origin instanceof MMapCachedDataSet) { + return (MMapCachedDataSet) origin; + } + return new MMapCachedDataSet(origin, defaultCacheDir()); + } + + /// @return the directory spill files are written to: `$DATASET_CACHE_DIR/mmap` or `dataset_cache/mmap` + public static Path defaultCacheDir() { + String env = System.getenv(ENV_DATASET_CACHE_DIR); + Path root = (env != null && !env.isEmpty()) ? Paths.get(env) : Paths.get("dataset_cache"); + return root.resolve("mmap"); + } + + /// Creates a mapped view of `origin`, spilling heap-resident base vectors to a file under `cacheDir`. + /// + /// @param origin the dataset to wrap + /// @param cacheDir where a spill file is written when the origin is not already memory-mapped + /// @throws UncheckedIOException if the spill file cannot be written or mapped + public MMapCachedDataSet(DataSet origin, Path cacheDir) { + this.origin = origin; + RandomAccessVectorValues source = origin.getBaseRavv(); + if (source instanceof MappedFvecsRandomAccessVectorValues) { + this.baseRavv = source; + logger.info("Serving {} base vectors of '{}' from mapped file {}", + source.size(), origin.getName(), ((MappedFvecsRandomAccessVectorValues) source).getPath()); + return; + } + Path file = cacheDir.resolve(safeFileName(origin.getName()) + "-" + source.size() + "x" + source.dimension() + ".fvecs"); + long expectedBytes = (long) source.size() * (Integer.BYTES + (long) source.dimension() * Float.BYTES); + long start = System.nanoTime(); + try { + Files.createDirectories(cacheDir); + boolean reused = Files.isRegularFile(file) && Files.size(file) == expectedBytes; + if (!reused) { + spill(source, file, cacheDir); + } + this.baseRavv = new MappedFvecsRandomAccessVectorValues(file); + logger.info("{} {} base vectors of '{}' {} {} and mapped them in {}s", + reused ? "Reused" : "Spilled", source.size(), origin.getName(), reused ? "from" : "to", file, + String.format("%.2f", (System.nanoTime() - start) / 1e9)); + } catch (IOException e) { + throw new UncheckedIOException("Failed to spill base vectors of '" + origin.getName() + "' to " + file, e); + } + } + + /// Writes `source` to a temporary file beside `file` and moves it into place. The target is never + /// rewritten in place: a spill file for the same dataset may already be mapped by another wrapper in + /// this JVM, and Windows refuses to modify a file with a live mapping. + private static void spill(RandomAccessVectorValues source, Path file, Path cacheDir) throws IOException { + Path temp = Files.createTempFile(cacheDir, file.getFileName().toString(), ".tmp"); + try { + SiftLoader.writeFvecs(temp, source); + try { + Files.move(temp, file, StandardCopyOption.ATOMIC_MOVE, StandardCopyOption.REPLACE_EXISTING); + } catch (AtomicMoveNotSupportedException e) { + Files.move(temp, file, StandardCopyOption.REPLACE_EXISTING); + } + } finally { + Files.deleteIfExists(temp); + } + } + + private static String safeFileName(String name) { + String safe = name.replaceAll("[^A-Za-z0-9._-]", "_"); + return safe.isEmpty() ? "dataset" : safe; + } + + @Override + public DataSet getOrigin() { + return origin; + } + + @Override + public RandomAccessVectorValues getBaseRavv() { + return baseRavv; + } + + @Override + public int getDimension() { + return baseRavv.dimension(); + } +} 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/SimpleDataSet.java index bf9c69376..0469f0377 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/SimpleDataSet.java @@ -23,21 +23,29 @@ import java.util.List; +/// A {@link DataSet} assembled from its parts: a base-vector {@link RandomAccessVectorValues} of any +/// kind, plus heap-resident query vectors and ground truth. public class SimpleDataSet implements DataSet { private final String name; private final VectorSimilarityFunction similarityFunction; - private final List> baseVectors; + private final RandomAccessVectorValues baseRavv; private final List> queryVectors; private final List> groundTruth; - private RandomAccessVectorValues baseRavv; + /// Creates a dataset over an arbitrary base-vector reader. + /// + /// @param name the dataset name + /// @param similarityFunction the similarity function the dataset was built for + /// @param baseRavv the base vectors; must be non-empty + /// @param queryVectors the query vectors; must be non-empty and match the base dimension + /// @param groundTruth one neighbor list per query vector public SimpleDataSet(String name, VectorSimilarityFunction similarityFunction, - List> baseVectors, + RandomAccessVectorValues baseRavv, List> queryVectors, List> groundTruth) { - if (baseVectors.isEmpty()) { + if (baseRavv.size() == 0) { throw new IllegalArgumentException("Base vectors must not be empty"); } if (queryVectors.isEmpty()) { @@ -47,7 +55,7 @@ public SimpleDataSet(String name, throw new IllegalArgumentException("Ground truth vectors must not be empty"); } - if (baseVectors.get(0).length() != queryVectors.get(0).length()) { + if (baseRavv.dimension() != queryVectors.get(0).length()) { throw new IllegalArgumentException("Base and query vectors must have the same dimensionality"); } if (queryVectors.size() != groundTruth.size()) { @@ -56,24 +64,44 @@ public SimpleDataSet(String name, this.name = name; this.similarityFunction = similarityFunction; - this.baseVectors = baseVectors; + this.baseRavv = baseRavv; 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()); + name, baseRavv.size(), queryVectors.size(), baseRavv.dimension()); + } + + /// Creates a dataset over heap-resident base vectors, served through a {@link ListRandomAccessVectorValues}. + /// + /// @param name the dataset name + /// @param similarityFunction the similarity function the dataset was built for + /// @param baseVectors the base vectors; must be non-empty + /// @param queryVectors the query vectors; must be non-empty and match the base dimension + /// @param groundTruth one neighbor list per query vector + public SimpleDataSet(String name, + VectorSimilarityFunction similarityFunction, + List> baseVectors, + List> queryVectors, + List> groundTruth) + { + this(name, similarityFunction, listRavv(baseVectors), queryVectors, groundTruth); + } + + private static RandomAccessVectorValues listRavv(List> baseVectors) { + if (baseVectors.isEmpty()) { + throw new IllegalArgumentException("Base vectors must not be empty"); + } + return new ListRandomAccessVectorValues(baseVectors, baseVectors.get(0).length()); } @Override public int getDimension() { - return getBaseVectors().get(0).length(); + return baseRavv.dimension(); } @Override public RandomAccessVectorValues getBaseRavv() { - if (baseRavv == null) { - baseRavv = new ListRandomAccessVectorValues(getBaseVectors(), getDimension()); - } return baseRavv; } @@ -87,11 +115,6 @@ public VectorSimilarityFunction getSimilarityFunction() { return similarityFunction; } - @Override - public List> getBaseVectors() { - return baseVectors; - } - @Override public List> getQueryVectors() { return queryVectors; 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..05a10ae86 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 @@ -120,7 +120,7 @@ public static Row fromDataSet(String datasetName, basePath, queryPath, groundTruthPath, - ds.getBaseVectors().size(), + ds.getBaseRavv().size(), ds.getQueryVectors().size(), ds.getGroundTruth().size(), ds.getDimension(), 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..598aa0cfd 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 @@ -18,41 +18,53 @@ import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; import io.github.jbellis.jvector.example.yaml.TestDataPartition; -import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import java.util.ArrayList; import java.util.List; /** - * Utility for partitioning a DataSet into multiple segments based on a distribution. + * Utility for partitioning a DataSet into multiple contiguous segments based on a distribution. + *

+ * Partitions are {@link RandomAccessVectorValues#range(int, int)} views over the base vectors, so no + * vectors are copied. Partition {@code i} covers global ordinals {@code [sum(sizes[0..i)), sum(sizes[0..i]))}, + * which is the ordering compaction relies on to map partition-local ordinals back to global ones. */ public final class DataSetPartitioner { private DataSetPartitioner() {} public static final class PartitionedData { - public final List>> vectors; + public final List vectors; public final List sizes; - public PartitionedData(List>> vectors, List sizes) { + public PartitionedData(List vectors, List sizes) { this.vectors = vectors; this.sizes = sizes; } } public static PartitionedData partition(DataSet ds, int numParts, TestDataPartition.Distribution distribution) { - return partition(ds.getBaseVectors(), numParts, distribution); + return partition(ds.getBaseRavv(), numParts, distribution); } - public static PartitionedData partition(List> baseVectors, int numParts, TestDataPartition.Distribution distribution) { + /** + * Splits {@code baseVectors} into {@code numParts} contiguous ranged views sized by {@code distribution}. + * + * @param baseVectors the vectors to partition + * @param numParts the number of partitions + * @param distribution how to size the partitions + * @return the partition views and their sizes, in partition order + */ + public static PartitionedData partition(RandomAccessVectorValues baseVectors, int numParts, TestDataPartition.Distribution distribution) { List sizes = distribution.computeSplitSizes(baseVectors.size(), numParts); - List>> parts = new ArrayList<>(numParts); + List parts = new ArrayList<>(numParts); int runningStart = 0; for (int size : sizes) { int start = runningStart; int end = start + size; runningStart = end; - parts.add(new ArrayList<>(baseVectors.subList(start, end))); + parts.add(baseVectors.range(start, end)); } return new PartitionedData(parts, sizes); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValues.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValues.java new file mode 100644 index 000000000..a2cdf745f --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValues.java @@ -0,0 +1,244 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; + +import java.util.Map; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.LongAdder; + +/// A thread-safe, bounded, least-recently-used cache of heap-resident vectors in front of any +/// {@link RandomAccessVectorValues}, organised in *grains* of `grainSize` consecutive ordinals. +/// +/// A read of ordinal `n` serves grain `n / grainSize` from the cache, loading it on a miss through +/// `origin.range(start, end).copy()` so the origin may be value-shared (a memory-mapped file, say). +/// At most `maxGrains` loaded grains are kept; when a load would exceed that, the grain with the +/// oldest last use is dropped. The bound is soft only while more than `maxGrains` grains are being +/// loaded concurrently, since a grain still loading is never evicted. +/// +/// Vectors handed out are independent heap objects owned by their grain. They stay valid after the +/// grain is evicted (the caller's reference keeps them alive), so this reader is not value-shared +/// and {@link #copy()} returns `this`: one cache is meant to be shared by every thread. +/// +/// Loads of distinct grains proceed in parallel; concurrent misses on the same grain wait for the +/// first loader. A failed load is retried by the next reader rather than poisoning the grain. +public final class LruGrainRandomAccessVectorValues implements RandomAccessVectorValues { + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + private final RandomAccessVectorValues origin; + private final int size; + private final int dimension; + private final int grainSize; + private final int maxGrains; + private final int grainCount; + + private final Map grains = new ConcurrentHashMap<>(); + private final Object evictionLock = new Object(); + private final AtomicLong clock = new AtomicLong(); + private final LongAdder hits = new LongAdder(); + private final LongAdder misses = new LongAdder(); + private final LongAdder evictions = new LongAdder(); + + private static final class Grain { + final CountDownLatch loaded = new CountDownLatch(1); + volatile VectorFloat[] vectors; + volatile long lastUsed; + } + + /// @param origin the vectors to cache; read through ranged views, so it may be value-shared + /// @param grainSize consecutive ordinals per grain; must be positive + /// @param maxGrains loaded grains to keep; must be positive + public LruGrainRandomAccessVectorValues(RandomAccessVectorValues origin, int grainSize, int maxGrains) { + if (grainSize <= 0) { + throw new IllegalArgumentException("grainSize must be positive, got " + grainSize); + } + if (maxGrains <= 0) { + throw new IllegalArgumentException("maxGrains must be positive, got " + maxGrains); + } + this.origin = origin; + this.size = origin.size(); + this.dimension = origin.dimension(); + this.grainSize = grainSize; + this.maxGrains = maxGrains; + this.grainCount = (int) (((long) size + grainSize - 1) / grainSize); + } + + /// @return consecutive ordinals per grain + public int grainSize() { + return grainSize; + } + + /// @return the maximum number of loaded grains kept resident + public int maxGrains() { + return maxGrains; + } + + /// @return the number of grains the origin divides into + public int grainCount() { + return grainCount; + } + + /// @return grains currently in the cache, loading ones included + public int residentGrains() { + return grains.size(); + } + + /// @return reads served from an already-loaded grain + public long hits() { + return hits.sum(); + } + + /// @return grain loads performed + public long misses() { + return misses.sum(); + } + + /// @return grains dropped to stay within {@link #maxGrains()} + public long evictions() { + return evictions.sum(); + } + + /// @return a one-line summary of the cache counters, for logs + public String stats() { + return String.format("grains=%d/%d resident=%d hits=%d misses=%d evictions=%d", + Math.min(grainCount, maxGrains), grainCount, grains.size(), hits(), misses(), evictions()); + } + + @Override + public int size() { + return size; + } + + @Override + public int dimension() { + return dimension; + } + + @Override + public VectorFloat getVector(int nodeId) { + int idx = Objects.checkIndex(nodeId, size); + int g = idx / grainSize; + Grain grain = grain(g); + return grain.vectors[idx - g * grainSize]; + } + + @Override + public void getVectorInto(int node, VectorFloat destinationVector, int offset) { + destinationVector.copyFrom(getVector(node), 0, offset, dimension); + } + + @Override + public boolean isValueShared() { + return false; + } + + @Override + public RandomAccessVectorValues copy() { + return this; + } + + private Grain grain(int g) { + while (true) { + Grain grain = grains.get(g); + if (grain == null) { + Grain fresh = new Grain(); + grain = grains.putIfAbsent(g, fresh); + if (grain == null) { + return load(g, fresh); + } + } + if (grain.vectors == null) { + awaitLoaded(grain); + if (grain.vectors == null) { + continue; // the loader failed and withdrew the grain; try again + } + } + hits.increment(); + grain.lastUsed = clock.incrementAndGet(); + return grain; + } + } + + private Grain load(int g, Grain grain) { + try { + grain.lastUsed = clock.incrementAndGet(); + evictIfNeeded(grain); + int start = g * grainSize; + int end = Math.min(size, start + grainSize); + RandomAccessVectorValues slice = origin.range(start, end).copy(); + VectorFloat[] vectors = new VectorFloat[end - start]; + for (int i = 0; i < vectors.length; i++) { + VectorFloat v = vts.createFloatVector(dimension); + slice.getVectorInto(i, v, 0); + vectors[i] = v; + } + grain.vectors = vectors; + misses.increment(); + grain.lastUsed = clock.incrementAndGet(); + return grain; + } catch (RuntimeException | Error e) { + grains.remove(g, grain); + throw e; + } finally { + grain.loaded.countDown(); + } + } + + private static void awaitLoaded(Grain grain) { + try { + grain.loaded.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("Interrupted while waiting for a grain to load", e); + } + } + + private void evictIfNeeded(Grain incoming) { + if (grains.size() <= maxGrains) { + return; + } + synchronized (evictionLock) { + while (grains.size() > maxGrains) { + Integer victimKey = null; + Grain victim = null; + for (var e : grains.entrySet()) { + Grain candidate = e.getValue(); + if (candidate == incoming || candidate.vectors == null) { + continue; // never evict a grain that is still loading + } + if (victim == null || candidate.lastUsed < victim.lastUsed) { + victim = candidate; + victimKey = e.getKey(); + } + } + if (victim == null) { + return; // everything resident is mid-load; the bound is exceeded until those finish + } + if (grains.remove(victimKey, victim)) { + evictions.increment(); + } + } + } + } +} diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValues.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValues.java new file mode 100644 index 000000000..57ffe0678 --- /dev/null +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValues.java @@ -0,0 +1,193 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; + +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.nio.MappedByteBuffer; +import java.nio.channels.FileChannel; +import java.nio.file.Path; +import java.nio.file.StandardOpenOption; +import java.util.Objects; + +/// A memory-mapped, read-only {@link RandomAccessVectorValues} over a standard `.fvecs` file. +/// +/// Each fvecs record is a little-endian `int` dimension followed by that many little-endian +/// `float`s. All records must have the same dimension, which is taken from the first record and +/// verified on every read. The file is mapped in slabs of at most {@link #DEFAULT_MAX_SLAB_BYTES} +/// so files larger than 2 GB work; a record never straddles two slabs. +/// +/// Reads are served by copying the record out of the mapping. Consequently {@link #getVector(int)} +/// returns a per-instance scratch vector and {@link #isValueShared()} is `true`; use +/// {@link #copy()} (or {@link #threadLocalSupplier()}) for concurrent readers, which share the +/// mapping and cost only a scratch vector each. Bulk in-memory caching should go through +/// {@link #range(int, int)} plus {@link #getVectorInto(int, VectorFloat, int)}. +/// +/// The mapping is released when the instance (and all copies) become unreachable; there is +/// deliberately no explicit unmap, since unmapping under a concurrent read crashes the JVM. +public final class MappedFvecsRandomAccessVectorValues implements RandomAccessVectorValues { + /// Upper bound on the bytes mapped per slab: 1 GiB, comfortably under the `int` limit of a single mapping. + public static final long DEFAULT_MAX_SLAB_BYTES = 1L << 30; + + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + private final Path path; + private final int dimension; + private final int size; + private final int stride; + private final int recordsPerSlab; + private final MappedByteBuffer[] slabs; + private final VectorFloat scratch; + + /// Maps the fvecs file at `path` using the default slab size. + /// + /// @param path the `.fvecs` file + /// @throws IOException if the file cannot be read, is empty, or is not a well-formed fvecs file + public MappedFvecsRandomAccessVectorValues(Path path) throws IOException { + this(path, DEFAULT_MAX_SLAB_BYTES); + } + + /// Maps the fvecs file at `path`, limiting each slab to `maxSlabBytes` (rounded down to a whole + /// number of records). Exposed for tests that need to exercise slab boundaries on small files. + /// + /// @param path the `.fvecs` file + /// @param maxSlabBytes the maximum bytes per mapped slab; must hold at least one record + /// @throws IOException if the file cannot be read, is empty, or is not a well-formed fvecs file + public MappedFvecsRandomAccessVectorValues(Path path, long maxSlabBytes) throws IOException { + this.path = path; + try (FileChannel channel = FileChannel.open(path, StandardOpenOption.READ)) { + long length = channel.size(); + if (length < Integer.BYTES) { + throw new IOException("fvecs file is empty or truncated: " + path); + } + ByteBuffer header = ByteBuffer.allocate(Integer.BYTES).order(ByteOrder.LITTLE_ENDIAN); + channel.read(header, 0); + int dim = header.getInt(0); + if (dim <= 0) { + throw new IOException("Corrupt fvecs file: negative or zero dimension " + dim + " (possible file corruption or wrong format): " + path); + } + if (dim > 100_000) { + throw new IOException("Unreasonable dimension " + dim + " in fvecs file (possible file corruption or wrong format): " + path); + } + this.dimension = dim; + this.stride = Integer.BYTES + dim * Float.BYTES; + if (length % stride != 0) { + throw new IOException("fvecs file length " + length + " is not a multiple of the " + stride + + "-byte record size for dimension " + dim + " (truncated or mixed dimensions): " + path); + } + long count = length / stride; + if (count > Integer.MAX_VALUE) { + throw new IOException("fvecs file holds " + count + " vectors, more than can be addressed by ordinal: " + path); + } + this.size = (int) count; + if (maxSlabBytes < stride) { + throw new IllegalArgumentException("maxSlabBytes " + maxSlabBytes + " is smaller than one record of " + stride + " bytes"); + } + this.recordsPerSlab = (int) Math.min(count, maxSlabBytes / stride); + int slabCount = (int) ((count + recordsPerSlab - 1) / recordsPerSlab); + this.slabs = new MappedByteBuffer[slabCount]; + long slabBytes = (long) recordsPerSlab * stride; + for (int s = 0; s < slabCount; s++) { + long offset = s * slabBytes; + long slabLength = Math.min(slabBytes, length - offset); + MappedByteBuffer slab = channel.map(FileChannel.MapMode.READ_ONLY, offset, slabLength); + slab.order(ByteOrder.LITTLE_ENDIAN); + slabs[s] = slab; + } + } + this.scratch = vts.createFloatVector(dimension); + } + + private MappedFvecsRandomAccessVectorValues(MappedFvecsRandomAccessVectorValues other) { + this.path = other.path; + this.dimension = other.dimension; + this.size = other.size; + this.stride = other.stride; + this.recordsPerSlab = other.recordsPerSlab; + this.slabs = other.slabs; + this.scratch = vts.createFloatVector(dimension); + } + + /// @return the mapped file + public Path getPath() { + return path; + } + + /// @return the number of mapped slabs backing this file + public int slabCount() { + return slabs.length; + } + + @Override + public int size() { + return size; + } + + @Override + public int dimension() { + return dimension; + } + + @Override + public VectorFloat getVector(int nodeId) { + read(nodeId, scratch, 0); + return scratch; + } + + @Override + public void getVectorInto(int node, VectorFloat destinationVector, int offset) { + read(node, destinationVector, offset); + } + + @Override + public boolean isValueShared() { + return true; + } + + /// Returns a reader that shares the file mapping but owns its own scratch vector. + @Override + public RandomAccessVectorValues copy() { + return new MappedFvecsRandomAccessVectorValues(this); + } + + private void read(int node, VectorFloat dest, int destOffset) { + int idx = Objects.checkIndex(node, size); + MappedByteBuffer slab = slabs[idx / recordsPerSlab]; + int pos = (idx % recordsPerSlab) * stride; + int recordDim = slab.getInt(pos); + if (recordDim != dimension) { + throw new IllegalStateException("Corrupt fvecs record " + idx + " in " + path + ": dimension " + recordDim + " != " + dimension); + } + int dataPos = pos + Integer.BYTES; + Object backing = dest.get(); + if (backing instanceof float[]) { + ByteBuffer view = slab.duplicate().order(ByteOrder.LITTLE_ENDIAN); + view.position(dataPos); + view.asFloatBuffer().get((float[]) backing, dest.offset(destOffset), dimension); + } else { + for (int i = 0; i < dimension; i++) { + dest.set(destOffset + i, slab.getFloat(dataPos + i * Float.BYTES)); + } + } + } +} 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..992d26adb 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 @@ -16,17 +16,21 @@ package io.github.jbellis.jvector.example.util; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import io.github.jbellis.jvector.vector.VectorizationProvider; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; import java.io.BufferedInputStream; +import java.io.BufferedOutputStream; import java.io.DataInputStream; import java.io.FileInputStream; import java.io.IOException; import java.io.UncheckedIOException; import java.nio.ByteBuffer; import java.nio.ByteOrder; +import java.nio.file.Files; +import java.nio.file.Path; import java.util.ArrayList; import java.util.HashSet; import java.util.List; @@ -60,6 +64,29 @@ public static List> readFvecs(String filePath) { return vectors; } + /// Writes every vector of `vectors` to `path` in fvecs format (little-endian `int` dimension + /// followed by the little-endian `float` components), overwriting any existing file. + /// + /// @param path the file to write + /// @param vectors the vectors to write, in ordinal order + /// @throws IOException if the file cannot be written + public static void writeFvecs(Path path, RandomAccessVectorValues vectors) throws IOException { + int dimension = vectors.dimension(); + var record = ByteBuffer.allocate(Integer.BYTES + dimension * Float.BYTES).order(ByteOrder.LITTLE_ENDIAN); + var scratch = vectorTypeSupport.createFloatVector(dimension); + try (var out = new BufferedOutputStream(Files.newOutputStream(path), 1 << 20)) { + for (int i = 0; i < vectors.size(); i++) { + vectors.getVectorInto(i, scratch, 0); + record.clear(); + record.putInt(dimension); + for (int d = 0; d < dimension; d++) { + record.putFloat(scratch.get(d)); + } + out.write(record.array()); + } + } + } + public static List> readIvecs(String filename) { var groundTruthTopK = new ArrayList>(); diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/DatasetCollection.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/DatasetCollection.java index fc1f7f351..eed1349dd 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/DatasetCollection.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/DatasetCollection.java @@ -16,16 +16,33 @@ package io.github.jbellis.jvector.example.yaml; +import io.github.jbellis.jvector.example.benchmarks.datasets.DataSetSpec; import org.yaml.snakeyaml.Yaml; -import software.amazon.awssdk.http.auth.aws.internal.signer.chunkedencoding.Chunk; import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; import java.util.ArrayList; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; +/// The named sections of `datasets.yml`, each a list of datasets to benchmark. +/// +/// A list entry is either a plain string in {@link DataSetSpec} sugared form, e.g. +/// `cohere-english-v3-100k` or `cohere-english-v3-100k(mmap)`, or a structured map: +/// ```yaml +/// regression-tests: +/// - cap-1M +/// - name: cohere-english-v3-1M +/// profile: default +/// wrappers: +/// - mmap +/// - lru: { grain: 4096, capacityMb: 512 } +/// ``` +/// Both forms are exposed as their canonical sugared string, which {@link +/// io.github.jbellis.jvector.example.benchmarks.datasets.DataSets} parses back into a spec, so +/// existing name-based filtering and configuration lookups keep working. public class DatasetCollection { private static final String defaultFile = "./jvector-examples/yaml-configs/datasets.yml"; @@ -40,9 +57,39 @@ public static DatasetCollection load() throws IOException { } public static DatasetCollection load(String file) throws IOException { - InputStream inputStream = new FileInputStream(file); - Yaml yaml = new Yaml(); - return new DatasetCollection(yaml.load(inputStream)); + try (InputStream inputStream = new FileInputStream(file)) { + Yaml yaml = new Yaml(); + Map> raw = yaml.load(inputStream); + return new DatasetCollection(canonicalize(raw)); + } + } + + /// Converts each section's entries to canonical sugared dataset spec strings. + /// + /// @param raw the parsed YAML: section name to list of string or map entries + /// @return section name to list of canonical spec strings; null sections are preserved as null + static Map> canonicalize(Map> raw) { + Map> result = new LinkedHashMap<>(); + if (raw == null) { + return result; + } + for (var section : raw.entrySet()) { + List entries = section.getValue(); + if (entries == null) { + result.put(section.getKey(), null); + continue; + } + List specs = new ArrayList<>(entries.size()); + for (Object entry : entries) { + try { + specs.add(DataSetSpec.from(entry).toString()); + } catch (IllegalArgumentException e) { + throw new IllegalArgumentException("Invalid dataset entry in section '" + section.getKey() + "': " + e.getMessage(), e); + } + } + result.put(section.getKey(), specs); + } + return result; } public List getAll() { diff --git a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/MultiConfig.java b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/MultiConfig.java index 6e56a6e26..324ea3bc5 100644 --- a/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/MultiConfig.java +++ b/jvector-examples/src/main/java/io/github/jbellis/jvector/example/yaml/MultiConfig.java @@ -16,6 +16,7 @@ package io.github.jbellis.jvector.example.yaml; +import io.github.jbellis.jvector.example.benchmarks.datasets.DataSetSpec; import io.github.jbellis.jvector.graph.disk.OnDiskGraphIndex; import org.yaml.snakeyaml.LoaderOptions; import org.yaml.snakeyaml.constructor.Constructor; @@ -50,8 +51,14 @@ public class MultiConfig { private static final java.util.concurrent.atomic.AtomicReference DEFAULT_FILE_USED = new java.util.concurrent.atomic.AtomicReference<>(); + /// Loads the per-dataset config from `index-parameters/.yml`, falling back to `default.yml`. + /// + /// `datasetName` may be a {@link DataSetSpec} in sugared form (e.g. `cap-1M(mmap)`): the config + /// file is looked up by the bare name, and when the spec carries a profile or wrappers the + /// resulting config's `dataset` is set to the full spec so that loading honors it. public static MultiConfig getDefaultConfig(String datasetName) throws FileNotFoundException { - var name = defaultDirectory + datasetName; + DataSetSpec spec = datasetName.endsWith(".yml") ? null : DataSetSpec.parse(datasetName); + var name = defaultDirectory + (spec == null ? datasetName : spec.getName()); if (!name.endsWith(".yml")) { name += ".yml"; } @@ -67,7 +74,7 @@ public static MultiConfig getDefaultConfig(String datasetName) throws FileNotFou var config = getConfig(configFile); - if (useDefault) { + if (useDefault || (spec != null && (spec.hasWrappers() || !spec.isDefaultProfile()))) { config.dataset = datasetName; } 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..8755e02f0 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 @@ -242,7 +242,8 @@ public static void benchmarkComparison(ImmutableGraphIndex graph, // Build suppliers for inline features (NVQ only - FUSED_ADC needs neighbors) Map> inlineSuppliers = new EnumMap<>(FeatureId.class); - inlineSuppliers.put(FeatureId.NVQ_VECTORS, ordinal -> new NVQ.State(nvq.encode(floatVectors.getVector(ordinal)))); + var vectors = floatVectors.threadLocalSupplier(); // suppliers run on parallel writer threads + inlineSuppliers.put(FeatureId.NVQ_VECTORS, ordinal -> new NVQ.State(nvq.encode(vectors.get().getVector(ordinal)))); // FUSED_ADC supplier needs graph view, provided at write time var identityMapper = new OrdinalMapper.IdentityMapper(floatVectors.size() - 1); @@ -305,7 +306,7 @@ public static void main(String[] args) throws IOException { DataSet ds = 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()); + System.out.printf("Loaded %d vectors of dimension %d%n", ds.getBaseRavv().size(), ds.getDimension()); var floatVectors = ds.getBaseRavv(); diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFDTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFDTest.java index 33379dd57..7c0f415fd 100644 --- a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFDTest.java +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetLoaderSimpleMFDTest.java @@ -69,6 +69,58 @@ public void setUp() throws IOException { testMetadata = DataSetMetadataReader.load(metadataFile.toString()); } + // ======================================================================== + // Profiles + // ======================================================================== + + @Test + public void profileKeysResolveThroughSpecs() throws IOException { + writeTestDataFiles(cacheDir); + Files.createDirectories(cacheDir.resolve("fast")); + writeLocalOverrideDataFiles(cacheDir.resolve("fast")); + Files.writeString(cacheDir.resolve("catalog_entries.yaml"), + "test-ds:\n" + + " base: test_base.fvecs\n" + + " query: test_query.fvecs\n" + + " gt: test_gt.ivecs\n" + + "\"test-ds:fast\":\n" + + " base: fast/test_base.fvecs\n" + + " query: fast/test_query.fvecs\n" + + " gt: fast/test_gt.ivecs\n" + + "\"sub-ds:default\":\n" + + " base: test_base.fvecs\n" + + " query: test_query.fvecs\n" + + " gt: test_gt.ivecs\n"); + var loader = new DataSetLoaderSimpleMFD(null, cacheDir.toString(), false, testMetadata); + + // default profile: bare entry + var plain = loader.loadDataSet(DataSetSpec.parse("test-ds")).orElseThrow(); + assertEquals("test-ds", plain.getName()); + assertEquals(5, plain.getDataSet().getBaseRavv().size()); + + // explicit profile: the name:profile entry, with metadata falling back to the bare name + var fast = loader.loadDataSet(DataSetSpec.parse("test-ds:fast")).orElseThrow(); + assertEquals("test-ds:fast", fast.getName()); + assertEquals(1, fast.getDataSet().getBaseRavv().size()); + + // default profile also matches a name:default entry when no bare entry exists + var sub = loader.loadDataSet(DataSetSpec.parse("sub-ds")).orElseThrow(); + assertEquals("sub-ds:default", sub.getName()); + assertEquals(5, sub.getDataSet().getBaseRavv().size()); + + // unknown profile is simply not found, not an error + assertTrue(loader.loadDataSet(DataSetSpec.parse("test-ds:slow")).isEmpty()); + + // the literal-key path still resolves the profile key as written + assertEquals("test-ds:fast", loader.loadDataSet("test-ds:fast").orElseThrow().getName()); + assertTrue(loader.loadDataSet("sub-ds").isEmpty()); + + // through the facade, the sugared name and wrappers compose with the profile + var ds = DataSets.loadDataSet("test-ds:fast(mmap)", java.util.List.of(loader)).orElseThrow().getDataSet(); + assertTrue(ds instanceof MMapCachedDataSet); + assertEquals(1, ds.getBaseRavv().size()); + } + // ======================================================================== // Basic loading // ======================================================================== @@ -87,7 +139,7 @@ public void loadsDatasetFromLocalCatalogAndFiles() throws IOException { assertEquals("test-ds", info.get().getName()); var ds = info.get().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); assertEquals(2, ds.getQueryVectors().size()); assertEquals(2, ds.getGroundTruth().size()); assertEquals(4, ds.getDimension()); @@ -200,7 +252,7 @@ public void loadsWithLocalPathAsYamlFile() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -243,7 +295,7 @@ public void nullCatalogUrlWorksWithLocalCatalog() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -256,7 +308,7 @@ public void emptyCatalogUrlWorksWithLocalCatalog() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -310,7 +362,7 @@ public void catalogUrlFetchesRemoteCatalogWhenNoLocalCatalogExists() throws IOEx assertTrue(Files.exists(cacheDir.resolve("catalog_entries.yaml"))); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); assertEquals(2, ds.getQueryVectors().size()); assertEquals(2, ds.getGroundTruth().size()); assertEquals(4, ds.getDimension()); @@ -437,7 +489,7 @@ public void subdirectoryDataFilesResolveRelativeToTheirCatalog() throws IOExcept // data files should resolve relative to subDir, not cacheDir var ds = loader.loadDataSet("sub-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -463,7 +515,7 @@ public void duplicateEntryAcrossCatalogsDoesNotFail() throws IOException { // should load without error — whichever catalog wins, the dataset is valid var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); assertNotNull(ds); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -569,7 +621,7 @@ public void base_urlOverrideIsUsedForDownload() throws IOException { ); var ds = loader.loadDataSet("private-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -611,7 +663,7 @@ public void base_urlWithoutTrailingSlashIsNormalized() throws IOException { // should load fine — base_url is normalized with trailing slash var ds = loader.loadDataSet("private-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -637,7 +689,7 @@ public void subdirectoryPathsInFileValuesResolveCorrectly() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); assertEquals(4, ds.getDimension()); } @@ -683,7 +735,7 @@ public void defaultsAreFoldedIntoEntries() throws IOException { // files exist locally so base_url isn't hit, but the entry should load fine var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -704,7 +756,7 @@ public void entryOverridesDefaults() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -752,7 +804,7 @@ public void cacheDirOverridesLocalDir() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -773,7 +825,7 @@ public void cacheDirFromDefaults() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -797,7 +849,7 @@ public void cacheDirEntryOverridesDefault() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -869,7 +921,7 @@ public void nonExistentCacheDirWithLocalFilesPrePopulated() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } // ======================================================================== @@ -894,7 +946,7 @@ public void envVarExpandedInBaseurl() throws IOException { // files exist locally so the expanded base_url isn't hit, but parsing should succeed var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -915,7 +967,7 @@ public void envVarExpandedInCacheDir() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -936,7 +988,7 @@ public void envVarExpandedInDefaults() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -970,7 +1022,7 @@ public void envVarWithDefaultUsesDefault() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -1047,7 +1099,7 @@ public void includeWithUnreachableRemoteWarnsButDoesNotFail() throws IOException // local entry should still work var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -1081,7 +1133,7 @@ public void includeWithMissingUrlFieldIsIgnored() throws IOException { ); var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -1104,7 +1156,7 @@ public void localEntryOverridesIncludedEntry() throws IOException { // local entry should work — the failed include shouldn't prevent it var ds = loader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, ds.getBaseVectors().size()); + assertEquals(5, ds.getBaseRavv().size()); } @Test @@ -1298,7 +1350,7 @@ public void includeOnlyCatalogLoadsOfflineFromCachedRemoteCatalog() throws IOExc ); var onlineDs = onlineLoader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, onlineDs.getBaseVectors().size()); + assertEquals(5, onlineDs.getBaseRavv().size()); assertTrue(Files.exists(cachedDataDir.resolve("test_base.fvecs"))); assertTrue(Files.exists(cachedDataDir.resolve("test_query.fvecs"))); assertTrue(Files.exists(cachedDataDir.resolve("test_gt.ivecs"))); @@ -1313,7 +1365,7 @@ public void includeOnlyCatalogLoadsOfflineFromCachedRemoteCatalog() throws IOExc ); var offlineDs = offlineLoader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(5, offlineDs.getBaseVectors().size()); + assertEquals(5, offlineDs.getBaseRavv().size()); assertEquals(2, offlineDs.getQueryVectors().size()); assertEquals(2, offlineDs.getGroundTruth().size()); assertEquals(4, offlineDs.getDimension()); @@ -1356,7 +1408,7 @@ public void localCatalogOverridesCachedIncludedRemoteCatalogOffline() throws IOE ); var onlineDs = onlineLoader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(1, onlineDs.getBaseVectors().size()); + assertEquals(1, onlineDs.getBaseRavv().size()); } finally { server.stop(0); } @@ -1367,7 +1419,7 @@ public void localCatalogOverridesCachedIncludedRemoteCatalogOffline() throws IOE ); var offlineDs = offlineLoader.loadDataSet("test-ds").orElseThrow().getDataSet(); - assertEquals(1, offlineDs.getBaseVectors().size()); + assertEquals(1, offlineDs.getBaseRavv().size()); assertEquals(1, offlineDs.getQueryVectors().size()); assertEquals(1, offlineDs.getGroundTruth().size()); assertEquals(4, offlineDs.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..9b7608a7b 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 @@ -350,7 +350,7 @@ public void dataSetInfoLazyLoading() { 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/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpecTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpecTest.java new file mode 100644 index 000000000..772fec3e3 --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetSpecTest.java @@ -0,0 +1,146 @@ +/* + * 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.example.benchmarks.datasets.DataSetSpec.WrapperSpec; +import org.junit.Test; + +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.*; + +/// Tests for {@link DataSetSpec} parsing of the sugared and structured forms, including wrapper options. +public class DataSetSpecTest { + + private static List names(DataSetSpec spec) { + return spec.getWrappers().stream().map(WrapperSpec::getName).collect(java.util.stream.Collectors.toList()); + } + + @Test + public void parsesNameOnly() { + var spec = DataSetSpec.parse("cohere-english-v3-100k"); + assertEquals("cohere-english-v3-100k", spec.getName()); + assertEquals(DataSetSpec.DEFAULT_PROFILE, spec.getProfile()); + assertTrue(spec.isDefaultProfile()); + assertFalse(spec.hasWrappers()); + assertEquals("cohere-english-v3-100k", spec.toString()); + } + + @Test + public void parsesProfileAndWrappers() { + var spec = DataSetSpec.parse("cohere:fast(memory, mmap)"); + assertEquals("cohere", spec.getName()); + assertEquals("fast", spec.getProfile()); + assertFalse(spec.isDefaultProfile()); + assertEquals(List.of("memory", "mmap"), names(spec)); + assertTrue(spec.getWrappers().get(0).getOptions().isEmpty()); + assertEquals("cohere:fast(memory,mmap)", spec.toString()); + + assertEquals(List.of("mmap"), names(DataSetSpec.parse("cohere(mmap)"))); + assertEquals("cohere(mmap)", DataSetSpec.parse("cohere(mmap)").toString()); + } + + @Test + public void parsesWrapperOptionsAndRoundTrips() { + var spec = DataSetSpec.parse("cohere(mmap, lru[grain=4096, capacityMb=512])"); + assertEquals(List.of("mmap", "lru"), names(spec)); + var lru = spec.getWrappers().get(1); + assertEquals(Map.of("grain", "4096", "capacityMb", "512"), lru.getOptions()); + assertEquals(List.of("grain", "capacityMb"), List.copyOf(lru.getOptions().keySet()), "option order is preserved"); + assertEquals("cohere(mmap,lru[grain=4096,capacityMb=512])", spec.toString()); + assertEquals(spec, DataSetSpec.parse(spec.toString())); + + assertEquals(new WrapperSpec("lru", null), DataSetSpec.parse("x(lru[])").getWrappers().get(0)); + assertEquals("x(lru)", DataSetSpec.parse("x(lru[ ])").toString()); + assertEquals("x(a[k=])", DataSetSpec.parse("x(a[k=])").toString()); + } + + @Test + public void defaultProfileIsCanonicalizedAway() { + var explicit = DataSetSpec.parse("cohere:default(mmap)"); + var implicit = DataSetSpec.parse("cohere(mmap)"); + assertEquals(implicit, explicit); + assertEquals(implicit.hashCode(), explicit.hashCode()); + assertEquals("cohere(mmap)", explicit.toString()); + assertEquals(DataSetSpec.DEFAULT_PROFILE, DataSetSpec.parse(" cohere : default ( ) ").getProfile()); + assertFalse(DataSetSpec.parse("cohere()").hasWrappers()); + } + + @Test + public void rejectsMalformedStrings() { + for (String bad : new String[] {"", " ", "a(b", "a:b:c", "a)b", "(mmap)", ":fast", "a(b)c", "a((b))", + "a(lru[grain)", "a(lru grain=1])", "a(lru[grain])", "a(lru[[x=1]])", "a(lru[x=1]])", "a[x=1]", "a(l r u)"}) { + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.parse(bad), "should reject '" + bad + "'"); + } + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.parse(null)); + assertThrows(IllegalArgumentException.class, () -> new DataSetSpec(" ", null, null)); + assertThrows(IllegalArgumentException.class, () -> new WrapperSpec(" ", null)); + assertThrows(IllegalArgumentException.class, () -> new WrapperSpec("lru", Map.of("grain", "1,2"))); + assertThrows(IllegalArgumentException.class, () -> new WrapperSpec("lru", Map.of("a=b", "1"))); + assertThrows(IllegalArgumentException.class, () -> new WrapperSpec("lru", Map.of(" ", "1"))); + } + + @Test + public void fromAcceptsStringsAndMaps() { + assertEquals(DataSetSpec.parse("cohere(mmap)"), DataSetSpec.from("cohere(mmap)")); + + var structured = DataSetSpec.from(Map.of("name", "cohere", "profile", "default", "wrappers", List.of("mmap"))); + assertEquals(DataSetSpec.parse("cohere(mmap)"), structured); + + var minimal = DataSetSpec.from(Map.of("name", "cohere")); + assertEquals(DataSetSpec.parse("cohere"), minimal); + + var stringWrappers = DataSetSpec.from(Map.of("name", "cohere", "profile", "fast", "wrappers", "memory,mmap")); + assertEquals(DataSetSpec.parse("cohere:fast(memory,mmap)"), stringWrappers); + + var same = DataSetSpec.parse("x"); + assertSame(same, DataSetSpec.from(same)); + } + + @Test + public void fromAcceptsWrapperOptionsInEveryForm() { + Map keyed = new LinkedHashMap<>(); + keyed.put("grain", 4096); + keyed.put("capacityMb", 512); + Map named = new LinkedHashMap<>(); + named.put("name", "lru"); + named.put("grain", 2048); + Map bare = new LinkedHashMap<>(); + bare.put("mmap", null); + + var spec = DataSetSpec.from(Map.of("name", "cohere", "wrappers", + List.of("memory", Map.of("lru", keyed), named, bare, "lru[grain=8]"))); + assertEquals("cohere(memory,lru[grain=4096,capacityMb=512],lru[grain=2048],mmap,lru[grain=8])", spec.toString()); + assertEquals(spec, DataSetSpec.parse(spec.toString())); + assertEquals("4096", spec.getWrappers().get(1).getOptions().get("grain")); + } + + @Test + public void fromRejectsBadMaps() { + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("profile", "fast"))); + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("name", "x", "loader", "y"))); + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("name", "x", "wrappers", 42))); + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(42)); + // a wrapper map with two keys and no name is ambiguous + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("name", "x", "wrappers", List.of(Map.of("lru", Map.of(), "mmap", Map.of()))))); + // options must be a map + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("name", "x", "wrappers", List.of(Map.of("lru", 4096))))); + assertThrows(IllegalArgumentException.class, () -> DataSetSpec.from(Map.of("name", "x", "wrappers", List.of(42)))); + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapperTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapperTest.java new file mode 100644 index 000000000..6c434044d --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/benchmarks/datasets/DataSetWrapperTest.java @@ -0,0 +1,358 @@ +/* + * 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.example.util.MappedFvecsRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorSimilarityFunction; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import org.junit.Rule; +import org.junit.Test; +import org.junit.rules.TemporaryFolder; + +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +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.Optional; +import java.util.concurrent.atomic.AtomicInteger; + +import static org.junit.jupiter.api.Assertions.*; + +/// Tests for the two built-in {@link DataSetWrapper}s and for how {@link DataSets} applies wrappers, +/// resolves symbolic wrapper names, and enforces the loader profile rule. +public class DataSetWrapperTest { + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + private static final int DIMENSION = 4; + + @Rule + public TemporaryFolder tempFolder = new TemporaryFolder(); + + private static final float[][] BASE = { + {1, 0, 0, 0}, {0, 1, 0, 0}, {0, 0, 1, 0}, {0, 0, 0, 1}, {0.5f, 0.5f, 0.5f, 0.5f}, {0.25f, 0.5f, 0.75f, 1}, + }; + + private static List> heapVectors(float[][] values) { + List> out = new ArrayList<>(); + for (float[] v : values) out.add(vts.createFloatVector(v.clone())); + return out; + } + + private static List> queries() { + return heapVectors(new float[][] {{1, 0, 0, 0}, {0, 0, 1, 0}}); + } + + private static List> groundTruth() { + return List.of(List.of(0, 4), List.of(2, 4)); + } + + private Path writeBaseFvecs() throws IOException { + Path file = tempFolder.newFile("base.fvecs").toPath(); + var buf = ByteBuffer.allocate(BASE.length * (Integer.BYTES + DIMENSION * Float.BYTES)).order(ByteOrder.LITTLE_ENDIAN); + for (float[] v : BASE) { + buf.putInt(DIMENSION); + for (float f : v) buf.putFloat(f); + } + Files.write(file, buf.array()); + return file; + } + + private DataSet mappedDataSet() throws IOException { + return new SimpleDataSet("mapped-ds", VectorSimilarityFunction.COSINE, + new MappedFvecsRandomAccessVectorValues(writeBaseFvecs()), queries(), groundTruth()); + } + + private static DataSet heapDataSet() { + return new SimpleDataSet("heap-ds", VectorSimilarityFunction.DOT_PRODUCT, heapVectors(BASE), queries(), groundTruth()); + } + + private static void assertBaseMatches(RandomAccessVectorValues ravv) { + assertEquals(BASE.length, ravv.size()); + assertEquals(DIMENSION, ravv.dimension()); + for (int i = 0; i < BASE.length; i++) { + VectorFloat v = ravv.getVector(i); + for (int d = 0; d < DIMENSION; d++) { + assertEquals(BASE[i][d], v.get(d), 0f); + } + } + } + + private static void assertDelegates(DataSetWrapper wrapper, DataSet origin) { + assertSame(origin, wrapper.getOrigin()); + assertEquals(origin.getName(), wrapper.getName()); + assertEquals(origin.getSimilarityFunction(), wrapper.getSimilarityFunction()); + assertSame(origin.getQueryVectors(), wrapper.getQueryVectors()); + assertSame(origin.getGroundTruth(), wrapper.getGroundTruth()); + assertEquals(origin.getDimension(), wrapper.getDimension()); + } + + // ------------------------------------------------------------------ InMemoryCachedDataSet + + @Test + public void inMemoryCacheCopiesMappedBaseVectorsToHeap() throws IOException { + DataSet origin = mappedDataSet(); + DataSetWrapper cached = InMemoryCachedDataSet.of(origin); + + assertTrue(cached instanceof InMemoryCachedDataSet); + assertDelegates(cached, origin); + assertTrue(cached.getBaseRavv() instanceof ListRandomAccessVectorValues); + assertFalse(cached.getBaseRavv().isValueShared()); + assertBaseMatches(cached.getBaseRavv()); + assertNotSame(cached.getBaseRavv().getVector(0), cached.getBaseRavv().getVector(1)); + + assertSame(cached, InMemoryCachedDataSet.of(cached)); + } + + @Test + public void inMemoryCacheAdoptsHeapResidentBaseVectors() { + DataSet origin = heapDataSet(); + DataSetWrapper cached = InMemoryCachedDataSet.of(origin); + assertSame(origin.getBaseRavv(), cached.getBaseRavv()); + assertDelegates(cached, origin); + } + + // ------------------------------------------------------------------ MMapCachedDataSet + + @Test + public void mmapCacheSpillsHeapBaseVectorsToFile() throws IOException { + DataSet origin = heapDataSet(); + Path cacheDir = tempFolder.getRoot().toPath().resolve("mmap-cache"); + var mapped = new MMapCachedDataSet(origin, cacheDir); + + assertDelegates(mapped, origin); + assertTrue(mapped.getBaseRavv() instanceof MappedFvecsRandomAccessVectorValues); + assertBaseMatches(mapped.getBaseRavv()); + Path spill = ((MappedFvecsRandomAccessVectorValues) mapped.getBaseRavv()).getPath(); + assertEquals(cacheDir.resolve("heap-ds-6x4.fvecs"), spill); + assertEquals(BASE.length * (Integer.BYTES + DIMENSION * Float.BYTES), Files.size(spill)); + + assertSame(mapped, MMapCachedDataSet.of(mapped)); + + // a second spill of the same dataset reuses the file, which is still mapped above, instead of rewriting it + var again = new MMapCachedDataSet(heapDataSet(), cacheDir); + assertEquals(spill, ((MappedFvecsRandomAccessVectorValues) again.getBaseRavv()).getPath()); + assertBaseMatches(again.getBaseRavv()); + } + + @Test + public void mmapCacheReplacesStaleSpillFile() throws IOException { + Path cacheDir = tempFolder.getRoot().toPath().resolve("stale"); + Files.createDirectories(cacheDir); + Path spill = cacheDir.resolve("heap-ds-6x4.fvecs"); + Files.write(spill, new byte[] {1, 2, 3}); // wrong size, so it cannot be a valid spill of this dataset + + var mapped = new MMapCachedDataSet(heapDataSet(), cacheDir); + assertEquals(spill, ((MappedFvecsRandomAccessVectorValues) mapped.getBaseRavv()).getPath()); + assertEquals(BASE.length * (Integer.BYTES + DIMENSION * Float.BYTES), Files.size(spill)); + assertBaseMatches(mapped.getBaseRavv()); + try (var files = Files.list(cacheDir)) { + assertEquals(1, files.count(), "temporary spill files are cleaned up"); + } + } + + @Test + public void mmapCacheAdoptsAlreadyMappedBaseVectors() throws IOException { + DataSet origin = mappedDataSet(); + Path cacheDir = tempFolder.getRoot().toPath().resolve("unused-cache"); + var mapped = new MMapCachedDataSet(origin, cacheDir); + assertSame(origin.getBaseRavv(), mapped.getBaseRavv()); + assertFalse(Files.exists(cacheDir)); + } + + @Test + public void wrappersCompose() throws IOException { + DataSet origin = heapDataSet(); + Path cacheDir = tempFolder.getRoot().toPath().resolve("compose"); + DataSetWrapper mapped = new MMapCachedDataSet(origin, cacheDir); + DataSetWrapper rehydrated = InMemoryCachedDataSet.of(mapped); + assertSame(mapped, rehydrated.getOrigin()); + assertTrue(rehydrated.getBaseRavv() instanceof ListRandomAccessVectorValues); + assertNotSame(origin.getBaseRavv(), rehydrated.getBaseRavv()); + assertBaseMatches(rehydrated.getBaseRavv()); + } + + // ------------------------------------------------------------------ LruCachedDataSet + + @Test + public void lruWrapperServesOriginThroughBoundedGrainCache() throws IOException { + DataSet origin = mappedDataSet(); + var lru = new LruCachedDataSet(origin, 2, 2L * DIMENSION * Float.BYTES * 2); // two grains of two vectors + assertDelegates(lru, origin); + assertEquals(2, lru.getBaseRavv().grainSize()); + assertEquals(2, lru.getBaseRavv().maxGrains()); + assertEquals(3, lru.getBaseRavv().grainCount()); + assertBaseMatches(lru.getBaseRavv()); + assertEquals(3, lru.getBaseRavv().misses()); + assertTrue(lru.getBaseRavv().evictions() >= 1); + assertFalse(lru.getBaseRavv().isValueShared()); + + assertSame(lru, LruCachedDataSet.of(lru)); + DataSetWrapper viaDefaults = LruCachedDataSet.of(origin); + assertTrue(viaDefaults instanceof LruCachedDataSet); + assertEquals(LruCachedDataSet.DEFAULT_GRAIN_SIZE, ((LruCachedDataSet) viaDefaults).getBaseRavv().grainSize()); + + DataSetWrapper tiny = LruCachedDataSet.provider(3, 1).wrap(origin); + assertEquals(1, ((LruCachedDataSet) tiny).getBaseRavv().maxGrains(), "capacity below one grain still keeps one grain"); + } + + // ------------------------------------------------------------------ DataSets facade + + /// A loader that serves {@code test-ds} from a heap dataset and counts materialisations. + private static class HeapLoader implements DataSetLoader { + final AtomicInteger loads = new AtomicInteger(); + + @Override + public Optional loadDataSet(String dataSetName) { + if (!dataSetName.equals("test-ds")) return Optional.empty(); + var props = new DataSetProperties.PropertyMap(Map.of( + DataSetProperties.KEY_NAME, "test-ds", + DataSetProperties.KEY_SIMILARITY_FUNCTION, VectorSimilarityFunction.DOT_PRODUCT, + DataSetProperties.KEY_LOAD_BEHAVIOR, DataSetProperties.LoadBehavior.NO_SCRUB)); + return Optional.of(new DataSetInfo(props, () -> { + loads.incrementAndGet(); + return heapDataSet(); + })); + } + } + + /// A loader that understands profiles and records the one it was asked for. + private static final class ProfileLoader extends HeapLoader { + String profileSeen; + + @Override + public Optional loadDataSet(DataSetSpec spec) { + profileSeen = spec.getProfile(); + return loadDataSet(spec.getName()); + } + } + + @Test + public void defaultWrappersCacheInMemoryLazily() { + var loader = new HeapLoader(); + var info = DataSets.loadDataSet("test-ds", List.of(loader)).orElseThrow(); + assertEquals("test-ds", info.getName()); + assertEquals(DataSetProperties.LoadBehavior.NO_SCRUB, info.loadBehavior()); + assertEquals(0, loader.loads.get(), "wrapping must not materialise the dataset"); + + DataSet ds = info.getDataSet(); + assertEquals(1, loader.loads.get()); + assertTrue(ds instanceof InMemoryCachedDataSet); + assertSame(ds, info.getDataSet()); + assertEquals(1, loader.loads.get()); + } + + @Test + public void symbolicWrappersReplaceTheDefaults() { + var loader = new HeapLoader(); + DataSet memory = DataSets.loadDataSet("test-ds(memory)", List.of(loader)).orElseThrow().getDataSet(); + assertTrue(memory instanceof InMemoryCachedDataSet); + + Path cacheDir = tempFolder.getRoot().toPath().resolve("facade"); + DataSetWrapper.Provider tempMmap = origin -> new MMapCachedDataSet(origin, cacheDir); + DataSets.wrapperProviders.put("mmap-tmp", DataSetWrapper.Factory.optionless("mmap-tmp", tempMmap)); + try { + DataSet mapped = DataSets.loadDataSet("test-ds:default(mmap-tmp)", List.of(loader)).orElseThrow().getDataSet(); + assertTrue(mapped instanceof MMapCachedDataSet); + assertTrue(mapped.getBaseRavv() instanceof MappedFvecsRandomAccessVectorValues); + + DataSet both = DataSets.loadDataSet("test-ds(mmap-tmp, memory)", List.of(loader)).orElseThrow().getDataSet(); + assertTrue(both instanceof InMemoryCachedDataSet); + assertTrue(((DataSetWrapper) both).getOrigin() instanceof MMapCachedDataSet); + assertBaseMatches(both.getBaseRavv()); + } finally { + DataSets.wrapperProviders.remove("mmap-tmp"); + } + + assertSame(MMapCachedDataSet.PROVIDER, DataSets.resolveWrappers(DataSetSpec.parse("x(mmap)").getWrappers()).get(0)); + assertSame(InMemoryCachedDataSet.PROVIDER, DataSets.resolveWrappers(DataSetSpec.parse("x(memory)").getWrappers()).get(0)); + assertSame(LruCachedDataSet.PROVIDER, DataSets.resolveWrappers(DataSetSpec.parse("x(lru)").getWrappers()).get(0)); + DataSet lru = DataSets.loadDataSet("test-ds(lru)", List.of(loader)).orElseThrow().getDataSet(); + assertTrue(lru instanceof LruCachedDataSet); + assertBaseMatches(lru.getBaseRavv()); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(bogus)", List.of(loader))); + } + + @Test + public void wrapperOptionsReachTheFactory() { + var loader = new HeapLoader(); + DataSet tuned = DataSets.loadDataSet("test-ds(lru[grain=2,capacityMb=1])", List.of(loader)).orElseThrow().getDataSet(); + var cache = ((LruCachedDataSet) tuned).getBaseRavv(); + assertEquals(2, cache.grainSize()); + assertEquals((1024 * 1024) / (2 * DIMENSION * Float.BYTES), cache.maxGrains()); + assertBaseMatches(cache); + + var structured = DataSetSpec.from(Map.of("name", "test-ds", "wrappers", List.of(Map.of("lru", Map.of("grain", 3))))); + DataSet fromYaml = DataSets.loadDataSet(structured, List.of(loader)).orElseThrow().getDataSet(); + assertEquals(3, ((LruCachedDataSet) fromYaml).getBaseRavv().grainSize()); + + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(lru[grain=0])", List.of(loader))); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(lru[grain=many])", List.of(loader))); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(lru[pages=4])", List.of(loader))); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(memory[grain=4])", List.of(loader))); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds(mmap[dir=x])", List.of(loader))); + assertSame(LruCachedDataSet.PROVIDER, LruCachedDataSet.provider(Map.of()), "no options yields the shared default provider"); + } + + @Test + public void explicitProvidersReplaceTheDefaults() { + var loader = new HeapLoader(); + DataSet raw = DataSets.loadDataSet("test-ds", List.of(loader), List.of()).orElseThrow().getDataSet(); + assertTrue(raw instanceof SimpleDataSet); + + var applied = new ArrayList(); + DataSetWrapper.Provider first = origin -> { applied.add("first"); return InMemoryCachedDataSet.of(origin); }; + DataSetWrapper.Provider second = origin -> { applied.add("second"); return InMemoryCachedDataSet.of(origin); }; + DataSet wrapped = DataSets.loadDataSet("test-ds", List.of(loader), List.of(first, second)).orElseThrow().getDataSet(); + assertEquals(List.of("first", "second"), applied); + assertTrue(wrapped instanceof InMemoryCachedDataSet); + + assertThrows(IllegalArgumentException.class, + () -> DataSets.loadDataSet("test-ds(memory)", List.of(loader), List.of(first))); + } + + @Test + public void profileRuleForLoadersWithoutProfileSupport() { + var loader = new HeapLoader(); + assertTrue(DataSets.loadDataSet("test-ds:default", List.of(loader)).isPresent()); + assertThrows(IllegalArgumentException.class, () -> DataSets.loadDataSet("test-ds:fast", List.of(loader))); + // a non-default profile for a dataset this loader does not have is not this loader's error + assertTrue(DataSets.loadDataSet("other-ds:fast", List.of(loader)).isEmpty()); + } + + @Test + public void profileAwareLoadersSeeTheProfile() { + var loader = new ProfileLoader(); + assertTrue(DataSets.loadDataSet("test-ds:fast(memory)", List.of(loader)).isPresent()); + assertEquals("fast", loader.profileSeen); + assertTrue(DataSets.loadDataSet("test-ds", List.of(loader)).isPresent()); + assertEquals(DataSetSpec.DEFAULT_PROFILE, loader.profileSeen); + } + + @Test + public void unknownDatasetIsEmptyAndHdf5NamesAreRejected() { + assertTrue(DataSets.loadDataSet("nope", List.of(new HeapLoader())).isEmpty()); + assertThrows(java.security.InvalidParameterException.class, + () -> DataSets.loadDataSet("test-ds.hdf5", List.of(new HeapLoader()))); + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/DataSetPartitionerTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/DataSetPartitionerTest.java new file mode 100644 index 000000000..4b99a147d --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/DataSetPartitionerTest.java @@ -0,0 +1,62 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.example.yaml.TestDataPartition; +import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RangeRandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import org.junit.Test; + +import java.util.ArrayList; +import java.util.List; + +import static org.junit.jupiter.api.Assertions.*; + +/// Tests that {@link DataSetPartitioner} produces contiguous, non-copying ranged views. +public class DataSetPartitionerTest { + + @Test + public void partitionsAreContiguousRangedViews() { + int count = 23, dimension = 2; + var vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + List> vectors = new ArrayList<>(); + for (int i = 0; i < count; i++) { + vectors.add(vts.createFloatVector(new float[] {i, -i})); + } + var base = new ListRandomAccessVectorValues(vectors, dimension); + + var parts = DataSetPartitioner.partition(base, 4, TestDataPartition.Distribution.UNIFORM); + assertEquals(4, parts.vectors.size()); + assertEquals(4, parts.sizes.size()); + assertEquals(count, parts.sizes.stream().mapToInt(Integer::intValue).sum()); + + int globalOrdinal = 0; + for (int p = 0; p < 4; p++) { + var view = parts.vectors.get(p); + assertEquals(parts.sizes.get(p).intValue(), view.size()); + assertTrue(view instanceof RangeRandomAccessVectorValues); + assertEquals(globalOrdinal, ((RangeRandomAccessVectorValues) view).fromOrdinal()); + for (int i = 0; i < view.size(); i++) { + assertSame(vectors.get(globalOrdinal), view.getVector(i)); + globalOrdinal++; + } + } + assertEquals(count, globalOrdinal); + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/FvecsLoadEconomyTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/FvecsLoadEconomyTest.java new file mode 100644 index 000000000..806ce464f --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/FvecsLoadEconomyTest.java @@ -0,0 +1,120 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.example.benchmarks.datasets.InMemoryCachedDataSet; +import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import org.junit.Test; + +import java.io.BufferedOutputStream; +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Random; + +import static org.junit.jupiter.api.Assertions.assertEquals; + +/// Compares the previous streaming fvecs loader against the memory-mapped reader plus in-memory +/// cache, both for load time and for post-load access time, and verifies they yield identical +/// vectors. Timings are reported, not asserted, since they depend on the host. +/// +/// The file size is controlled by `-Dfvecs.economy.vectors` and `-Dfvecs.economy.dim`; the file is +/// generated under `target/fvecs-economy` and reused across runs when its size matches. +public class FvecsLoadEconomyTest { + + @Test + public void mappedPathMatchesStreamingLoaderAndReportsTiming() throws IOException { + int count = Integer.getInteger("fvecs.economy.vectors", 20_000); + int dimension = Integer.getInteger("fvecs.economy.dim", 256); + int rounds = Integer.getInteger("fvecs.economy.rounds", 3); + Path dir = Path.of(System.getProperty("basedir", "jvector-examples"), "target", "fvecs-economy"); + Files.createDirectories(dir); + Path file = dir.resolve("economy-" + count + "x" + dimension + ".fvecs"); + long expectedBytes = (long) count * (Integer.BYTES + dimension * Float.BYTES); + if (!Files.exists(file) || Files.size(file) != expectedBytes) { + writeRandomFvecs(file, count, dimension); + } + System.out.printf("fvecs economy: %d vectors x %d dims (%.1f MB), %d rounds, vector provider %s%n", + count, dimension, expectedBytes / (1024.0 * 1024.0), rounds, + VectorizationProvider.getInstance().getClass().getSimpleName()); + System.out.printf("%-8s %14s %14s %14s %14s %14s%n", + "round", "stream-load", "map+cache", "map-only", "scan-stream", "scan-cached"); + + List> streamed = null; + List> cached = null; + for (int round = 0; round < rounds; round++) { + long t0 = System.nanoTime(); + streamed = SiftLoader.readFvecs(file.toString()); + long streamLoad = System.nanoTime() - t0; + + t0 = System.nanoTime(); + var mapped = new MappedFvecsRandomAccessVectorValues(file); + long mapOnly = System.nanoTime() - t0; + cached = InMemoryCachedDataSet.readAllVectors(mapped); + long mapAndCache = System.nanoTime() - t0; + + var streamedRavv = new ListRandomAccessVectorValues(streamed, dimension); + var cachedRavv = new ListRandomAccessVectorValues(cached, dimension); + long scanStream = timeScan(streamedRavv); + long scanCached = timeScan(cachedRavv); + + System.out.printf("%-8d %12.1fms %12.1fms %12.1fms %12.1fms %12.1fms%n", + round, streamLoad / 1e6, mapAndCache / 1e6, mapOnly / 1e6, scanStream / 1e6, scanCached / 1e6); + } + + assertEquals(streamed.size(), cached.size()); + for (int i = 0; i < count; i++) { + VectorFloat a = streamed.get(i); + VectorFloat b = cached.get(i); + for (int d = 0; d < dimension; d++) { + assertEquals(a.get(d), b.get(d), 0f, "vector " + i + " component " + d); + } + } + } + + private static long timeScan(RandomAccessVectorValues ravv) { + long t0 = System.nanoTime(); + float sink = 0; + for (int i = 0; i < ravv.size(); i++) { + VectorFloat v = ravv.getVector(i); + sink += v.get(0) + v.get(v.length() - 1); + } + long elapsed = System.nanoTime() - t0; + if (Float.isNaN(sink)) System.out.println("unexpected NaN"); + return elapsed; + } + + private static void writeRandomFvecs(Path file, int count, int dimension) throws IOException { + var random = new Random(1234); + var record = ByteBuffer.allocate(Integer.BYTES + dimension * Float.BYTES).order(ByteOrder.LITTLE_ENDIAN); + try (var out = new BufferedOutputStream(Files.newOutputStream(file), 1 << 20)) { + for (int i = 0; i < count; i++) { + record.clear(); + record.putInt(dimension); + for (int d = 0; d < dimension; d++) record.putFloat(random.nextFloat()); + out.write(record.array()); + } + } + VectorizationProvider.getInstance(); // keep the provider warm for the timed runs + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LargerThanHeapDataSetTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LargerThanHeapDataSetTest.java new file mode 100644 index 000000000..2835a6efe --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LargerThanHeapDataSetTest.java @@ -0,0 +1,124 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.LruCachedDataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.MMapCachedDataSet; +import io.github.jbellis.jvector.example.benchmarks.datasets.SimpleDataSet; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorSimilarityFunction; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import org.junit.Test; + +import java.io.BufferedOutputStream; +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Random; +import java.util.stream.IntStream; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/// Reads a whole fvecs dataset through the `mmap` and `lru` wrappers and checks both against a +/// direct mapped scan, without ever holding the base vectors in heap memory. +/// +/// Sized by `-Dfvecs.ltm.vectors` and `-Dfvecs.ltm.dim` (the file is shared with +/// {@link FvecsLoadEconomyTest} under `target/fvecs-economy`). To demonstrate larger-than-heap +/// operation, run with a file well above the heap, e.g. +/// `-Dfvecs.ltm.vectors=1000000 -Dfvecs.ltm.dim=1024 -DargLine=-Xmx1g`. +public class LargerThanHeapDataSetTest { + + @Test + public void mmapAndLruWrappersScanWithoutHeapResidentBaseVectors() throws IOException { + int count = Integer.getInteger("fvecs.ltm.vectors", 20_000); + int dimension = Integer.getInteger("fvecs.ltm.dim", 256); + long capacityBytes = Long.getLong("fvecs.ltm.lruCapacityMb", 64L) * 1024 * 1024; + Path dir = Path.of(System.getProperty("basedir", "jvector-examples"), "target", "fvecs-economy"); + Files.createDirectories(dir); + Path file = dir.resolve("economy-" + count + "x" + dimension + ".fvecs"); + long expectedBytes = (long) count * (Integer.BYTES + dimension * Float.BYTES); + if (!Files.exists(file) || Files.size(file) != expectedBytes) { + writeRandomFvecs(file, count, dimension); + } + long maxHeap = Runtime.getRuntime().maxMemory(); + System.out.printf("larger-than-heap: %.1f MB of vectors, max heap %.1f MB, lru capacity %.1f MB, vector provider %s%n", + expectedBytes / (1024.0 * 1024.0), maxHeap / (1024.0 * 1024.0), capacityBytes / (1024.0 * 1024.0), + VectorizationProvider.getInstance().getClass().getSimpleName()); + + var vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + var query = vts.createFloatVector(dimension); + DataSet origin = new SimpleDataSet("ltm", VectorSimilarityFunction.EUCLIDEAN, + new MappedFvecsRandomAccessVectorValues(file), List.of(query), List.of(List.of(0))); + + long t0 = System.nanoTime(); + double direct = checksum(new MappedFvecsRandomAccessVectorValues(file), 1); + long directMs = (System.nanoTime() - t0) / 1_000_000; + + DataSet mmap = MMapCachedDataSet.of(origin); + t0 = System.nanoTime(); + double viaMmap = checksum(mmap.getBaseRavv(), 8); + long mmapMs = (System.nanoTime() - t0) / 1_000_000; + + var lru = new LruCachedDataSet(origin, LruCachedDataSet.DEFAULT_GRAIN_SIZE, capacityBytes); + t0 = System.nanoTime(); + double viaLru = checksum(lru.getBaseRavv(), 8); + long lruMs = (System.nanoTime() - t0) / 1_000_000; + + System.out.printf("direct scan %d ms, mmap wrapper (8 threads) %d ms, lru wrapper (8 threads) %d ms; lru %s%n", + directMs, mmapMs, lruMs, lru.getBaseRavv().stats()); + + assertEquals(direct, viaMmap, 0.0); + assertEquals(direct, viaLru, 0.0); + assertTrue(lru.getBaseRavv().residentGrains() <= lru.getBaseRavv().maxGrains() + 8, lru.getBaseRavv().stats()); + assertTrue(expectedBytes > capacityBytes || count < 100_000, "test parameters should exceed the lru capacity"); + } + + /// Order-independent checksum over first, middle and last component of every vector. + private static double checksum(RandomAccessVectorValues ravv, int threads) { + int n = ravv.size(); + int chunk = (n + threads - 1) / threads; + return IntStream.range(0, threads).parallel().mapToDouble(t -> { + RandomAccessVectorValues local = ravv.copy(); + double sum = 0; + int end = Math.min(n, (t + 1) * chunk); + for (int i = t * chunk; i < end; i++) { + VectorFloat v = local.getVector(i); + sum += v.get(0) + v.get(v.length() / 2) + v.get(v.length() - 1); + } + return sum; + }).sum(); + } + + private static void writeRandomFvecs(Path file, int count, int dimension) throws IOException { + var random = new Random(1234); + var record = ByteBuffer.allocate(Integer.BYTES + dimension * Float.BYTES).order(ByteOrder.LITTLE_ENDIAN); + try (var out = new BufferedOutputStream(Files.newOutputStream(file), 1 << 20)) { + for (int i = 0; i < count; i++) { + record.clear(); + record.putInt(dimension); + for (int d = 0; d < dimension; d++) record.putFloat(random.nextFloat()); + out.write(record.array()); + } + } + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValuesTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValuesTest.java new file mode 100644 index 000000000..0c7efbf4a --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/LruGrainRandomAccessVectorValuesTest.java @@ -0,0 +1,195 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import org.junit.Rule; +import org.junit.Test; +import org.junit.rules.TemporaryFolder; + +import java.io.IOException; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.Random; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; + +import static io.github.jbellis.jvector.example.util.MappedFvecsRandomAccessVectorValuesTest.assertVectorEquals; +import static io.github.jbellis.jvector.example.util.MappedFvecsRandomAccessVectorValuesTest.randomVectors; +import static io.github.jbellis.jvector.example.util.MappedFvecsRandomAccessVectorValuesTest.writeFvecs; +import static org.junit.jupiter.api.Assertions.*; + +/// Tests for {@link LruGrainRandomAccessVectorValues}: correctness across grain and slab boundaries, +/// eviction bookkeeping, concurrent access, and recovery from a failed grain load. +public class LruGrainRandomAccessVectorValuesTest { + + @Rule + public TemporaryFolder tempFolder = new TemporaryFolder(); + + private MappedFvecsRandomAccessVectorValues mapped(float[][] vectors) throws IOException { + Path file = tempFolder.newFile("base.fvecs").toPath(); + writeFvecs(file, vectors, null); + int stride = Integer.BYTES + vectors[0].length * Float.BYTES; + return new MappedFvecsRandomAccessVectorValues(file, 5L * stride); // slabs of 5, deliberately misaligned with grains + } + + @Test + public void servesCorrectVectorsWithinTheGrainBound() throws IOException { + float[][] expected = randomVectors(100, 3, 1); + var cache = new LruGrainRandomAccessVectorValues(mapped(expected), 7, 3); + assertEquals(15, cache.grainCount()); + assertEquals(100, cache.size()); + assertEquals(3, cache.dimension()); + assertFalse(cache.isValueShared()); + assertSame(cache, cache.copy()); + + for (int i = 0; i < 100; i++) { + assertVectorEquals(expected[i], cache.getVector(i), 0); + } + assertEquals(15, cache.misses()); + assertEquals(100 - 15, cache.hits()); + assertEquals(12, cache.evictions()); + assertTrue(cache.residentGrains() <= 3, cache.stats()); + + // reading backwards re-faults grains the forward pass evicted + for (int i = 99; i >= 0; i--) { + assertVectorEquals(expected[i], cache.getVector(i), 0); + } + assertTrue(cache.misses() > 15, cache.stats()); + assertTrue(cache.residentGrains() <= 3, cache.stats()); + assertThrows(IndexOutOfBoundsException.class, () -> cache.getVector(100)); + } + + @Test + public void handedOutVectorsSurviveEviction() throws IOException { + float[][] expected = randomVectors(40, 4, 2); + var cache = new LruGrainRandomAccessVectorValues(mapped(expected), 4, 2); + VectorFloat first = cache.getVector(0); + for (int i = 0; i < 40; i++) { + cache.getVector(i); + } + assertTrue(cache.evictions() > 0); + assertVectorEquals(expected[0], first, 0); + // a re-faulted grain yields fresh objects, not the evicted ones + assertNotSame(first, cache.getVector(0)); + assertVectorEquals(expected[0], cache.getVector(0), 0); + } + + @Test + public void rangeViewsAndGetVectorIntoWork() throws IOException { + float[][] expected = randomVectors(30, 2, 3); + var cache = new LruGrainRandomAccessVectorValues(mapped(expected), 8, 2); + RandomAccessVectorValues view = cache.range(10, 25); + assertFalse(view.isValueShared()); + var vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + for (int i = 0; i < view.size(); i++) { + assertVectorEquals(expected[10 + i], view.getVector(i), 0); + var dest = vts.createFloatVector(4); + view.getVectorInto(i, dest, 2); + assertVectorEquals(expected[10 + i], dest, 2); + } + } + + @Test + public void wholeDatasetFitsWhenCapacityAllows() throws IOException { + float[][] expected = randomVectors(50, 2, 4); + var cache = new LruGrainRandomAccessVectorValues(mapped(expected), 10, 100); + for (int pass = 0; pass < 3; pass++) { + for (int i = 0; i < 50; i++) { + assertVectorEquals(expected[i], cache.getVector(i), 0); + } + } + assertEquals(5, cache.misses()); + assertEquals(0, cache.evictions()); + assertEquals(5, cache.residentGrains()); + } + + @Test + public void concurrentReadersSeeConsistentValues() throws Exception { + float[][] expected = randomVectors(5000, 8, 5); + var cache = new LruGrainRandomAccessVectorValues(mapped(expected), 64, 6); + int threads = 16; + ExecutorService pool = Executors.newFixedThreadPool(threads); + try { + List> futures = new ArrayList<>(); + for (int t = 0; t < threads; t++) { + long seed = t; + futures.add(pool.submit(() -> { + var random = new Random(seed); + for (int k = 0; k < 20_000; k++) { + int i = random.nextInt(expected.length); + assertVectorEquals(expected[i], cache.getVector(i), 0); + } + })); + } + for (Future f : futures) { + f.get(2, TimeUnit.MINUTES); + } + } finally { + pool.shutdownNow(); + } + assertTrue(cache.evictions() > 0, cache.stats()); + assertTrue(cache.residentGrains() <= 6 + threads, cache.stats()); + // every read is counted exactly once, as a hit or as the load that satisfied it + assertEquals(threads * 20_000L, cache.hits() + cache.misses(), cache.stats()); + } + + @Test + public void failedLoadIsRetriedByTheNextReader() { + var vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + List> vectors = new ArrayList<>(); + for (int i = 0; i < 10; i++) { + vectors.add(vts.createFloatVector(new float[] {i, i})); + } + var backing = new ListRandomAccessVectorValues(vectors, 2); + var failOnce = new AtomicBoolean(true); + RandomAccessVectorValues flaky = new RandomAccessVectorValues() { + @Override public int size() { return backing.size(); } + @Override public int dimension() { return backing.dimension(); } + @Override public VectorFloat getVector(int nodeId) { + if (nodeId >= 5 && failOnce.compareAndSet(true, false)) { + throw new IllegalStateException("transient read failure"); + } + return backing.getVector(nodeId); + } + @Override public boolean isValueShared() { return false; } + @Override public RandomAccessVectorValues copy() { return this; } + }; + + var cache = new LruGrainRandomAccessVectorValues(flaky, 5, 2); + assertEquals(1f, cache.getVector(1).get(0), 0f); + assertThrows(IllegalStateException.class, () -> cache.getVector(7)); + assertEquals(1, cache.residentGrains()); + assertEquals(7f, cache.getVector(7).get(0), 0f); + assertEquals(2, cache.residentGrains()); + assertEquals(2, cache.misses()); + } + + @Test + public void rejectsBadParameters() throws IOException { + var origin = mapped(randomVectors(3, 2, 6)); + assertThrows(IllegalArgumentException.class, () -> new LruGrainRandomAccessVectorValues(origin, 0, 1)); + assertThrows(IllegalArgumentException.class, () -> new LruGrainRandomAccessVectorValues(origin, 1, 0)); + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValuesTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValuesTest.java new file mode 100644 index 000000000..d1065a225 --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/util/MappedFvecsRandomAccessVectorValuesTest.java @@ -0,0 +1,215 @@ +/* + * 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.util; + +import io.github.jbellis.jvector.example.benchmarks.datasets.InMemoryCachedDataSet; +import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +import io.github.jbellis.jvector.vector.types.VectorTypeSupport; +import org.junit.Rule; +import org.junit.Test; +import org.junit.rules.TemporaryFolder; + +import java.io.IOException; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Random; + +import static org.junit.jupiter.api.Assertions.*; + +/// Tests for {@link MappedFvecsRandomAccessVectorValues}, including slab boundaries on small files. +public class MappedFvecsRandomAccessVectorValuesTest { + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + + @Rule + public TemporaryFolder tempFolder = new TemporaryFolder(); + + static float[][] randomVectors(int count, int dimension, long seed) { + var random = new Random(seed); + float[][] vectors = new float[count][dimension]; + for (float[] v : vectors) { + for (int d = 0; d < dimension; d++) { + v[d] = random.nextFloat() * 2 - 1; + } + } + return vectors; + } + + static void writeFvecs(Path path, float[][] vectors, int[] headerDims) throws IOException { + int dimension = vectors[0].length; + int bytesPerVector = Integer.BYTES + dimension * Float.BYTES; + var buf = ByteBuffer.allocate(vectors.length * bytesPerVector).order(ByteOrder.LITTLE_ENDIAN); + for (int i = 0; i < vectors.length; i++) { + buf.putInt(headerDims == null ? dimension : headerDims[i]); + for (float v : vectors[i]) buf.putFloat(v); + } + Files.write(path, buf.array()); + } + + static void assertVectorEquals(float[] expected, VectorFloat actual, int offset) { + for (int d = 0; d < expected.length; d++) { + assertEquals(expected[d], actual.get(offset + d), 0f, "component " + d); + } + } + + @Test + public void matchesStreamingReaderAcrossSlabBoundaries() throws IOException { + int count = 1000, dimension = 7; + float[][] expected = randomVectors(count, dimension, 42); + Path file = tempFolder.newFile("base.fvecs").toPath(); + writeFvecs(file, expected, null); + + int stride = Integer.BYTES + dimension * Float.BYTES; + var mapped = new MappedFvecsRandomAccessVectorValues(file, 5L * stride + 3); // 5 records per slab + assertEquals(count, mapped.size()); + assertEquals(dimension, mapped.dimension()); + assertEquals(200, mapped.slabCount()); + assertTrue(mapped.isValueShared()); + assertEquals(file, mapped.getPath()); + + List> streamed = SiftLoader.readFvecs(file.toString()); + assertEquals(count, streamed.size()); + for (int i = 0; i < count; i++) { + assertVectorEquals(expected[i], mapped.getVector(i), 0); + assertVectorEquals(expected[i], streamed.get(i), 0); + + var dest = vts.createFloatVector(2 * dimension); + mapped.getVectorInto(i, dest, dimension); + assertVectorEquals(expected[i], dest, dimension); + } + + // a default-sized mapping of the same file is a single slab with identical contents + var single = new MappedFvecsRandomAccessVectorValues(file); + assertEquals(1, single.slabCount()); + for (int i = 0; i < count; i += 97) { + assertVectorEquals(expected[i], single.getVector(i), 0); + } + } + + @Test + public void rangeViewsReadRebasedOrdinals() throws IOException { + float[][] expected = randomVectors(50, 3, 7); + Path file = tempFolder.newFile("base.fvecs").toPath(); + writeFvecs(file, expected, null); + var mapped = new MappedFvecsRandomAccessVectorValues(file, 4L * (Integer.BYTES + 3 * Float.BYTES)); + + var view = mapped.range(17, 31); + assertEquals(14, view.size()); + assertTrue(view.isValueShared()); + for (int i = 0; i < view.size(); i++) { + assertVectorEquals(expected[17 + i], view.getVector(i), 0); + } + assertThrows(IndexOutOfBoundsException.class, () -> view.getVector(14)); + } + + @Test + public void copiesShareTheMappingButNotTheScratchVector() throws IOException { + float[][] expected = randomVectors(4, 5, 11); + Path file = tempFolder.newFile("base.fvecs").toPath(); + writeFvecs(file, expected, null); + var a = new MappedFvecsRandomAccessVectorValues(file); + var b = a.copy(); + assertNotSame(a, b); + assertTrue(b instanceof MappedFvecsRandomAccessVectorValues); + assertEquals(a.slabCount(), ((MappedFvecsRandomAccessVectorValues) b).slabCount()); + + VectorFloat fromA = a.getVector(0); + VectorFloat fromB = b.getVector(3); + assertNotSame(fromA, fromB); + assertVectorEquals(expected[0], fromA, 0); + assertVectorEquals(expected[3], fromB, 0); + + // the shared scratch is overwritten by the next read on the same instance + assertSame(fromA, a.getVector(1)); + assertVectorEquals(expected[1], fromA, 0); + } + + @Test + public void parallelCachingThroughRangesMatchesFile() throws IOException { + int count = 3000, dimension = 16; + float[][] expected = randomVectors(count, dimension, 99); + Path file = tempFolder.newFile("base.fvecs").toPath(); + writeFvecs(file, expected, null); + var mapped = new MappedFvecsRandomAccessVectorValues(file, 64L * (Integer.BYTES + dimension * Float.BYTES)); + + List> cached = InMemoryCachedDataSet.readAllVectors(mapped); + assertEquals(count, cached.size()); + for (int i = 0; i < count; i++) { + assertVectorEquals(expected[i], cached.get(i), 0); + } + // every cached vector is an independent object + assertNotSame(cached.get(0), cached.get(1)); + var ravv = new ListRandomAccessVectorValues(cached, dimension); + assertFalse(ravv.isValueShared()); + } + + @Test + public void rejectsMalformedFiles() throws IOException { + Path empty = tempFolder.newFile("empty.fvecs").toPath(); + assertThrows(IOException.class, () -> new MappedFvecsRandomAccessVectorValues(empty)); + + Path zeroDim = tempFolder.newFile("zero.fvecs").toPath(); + Files.write(zeroDim, ByteBuffer.allocate(8).order(ByteOrder.LITTLE_ENDIAN).putInt(0).putFloat(1f).array()); + assertThrows(IOException.class, () -> new MappedFvecsRandomAccessVectorValues(zeroDim)); + + Path truncated = tempFolder.newFile("truncated.fvecs").toPath(); + var buf = ByteBuffer.allocate(Integer.BYTES + 2 * Float.BYTES + 3).order(ByteOrder.LITTLE_ENDIAN); + buf.putInt(2).putFloat(1f).putFloat(2f).put((byte) 1).put((byte) 2).put((byte) 3); + Files.write(truncated, buf.array()); + assertThrows(IOException.class, () -> new MappedFvecsRandomAccessVectorValues(truncated)); + + Path ok = tempFolder.newFile("ok.fvecs").toPath(); + writeFvecs(ok, randomVectors(3, 2, 1), null); + assertThrows(IllegalArgumentException.class, () -> new MappedFvecsRandomAccessVectorValues(ok, 4)); + } + + @Test + public void detectsCorruptRecordHeaderOnRead() throws IOException { + float[][] vectors = randomVectors(3, 4, 5); + Path file = tempFolder.newFile("corrupt.fvecs").toPath(); + writeFvecs(file, vectors, new int[] {4, 9, 4}); + var mapped = new MappedFvecsRandomAccessVectorValues(file); + assertEquals(3, mapped.size()); + assertVectorEquals(vectors[0], mapped.getVector(0), 0); + assertVectorEquals(vectors[2], mapped.getVector(2), 0); + assertThrows(IllegalStateException.class, () -> mapped.getVector(1)); + } + + @Test + public void writeFvecsRoundTrips() throws IOException { + float[][] expected = randomVectors(20, 6, 3); + List> vectors = new java.util.ArrayList<>(); + for (float[] v : expected) vectors.add(vts.createFloatVector(v)); + Path file = tempFolder.getRoot().toPath().resolve("written.fvecs"); + SiftLoader.writeFvecs(file, new ListRandomAccessVectorValues(vectors, 6)); + var mapped = new MappedFvecsRandomAccessVectorValues(file); + assertEquals(20, mapped.size()); + for (int i = 0; i < 20; i++) { + assertVectorEquals(expected[i], mapped.getVector(i), 0); + } + // overwriting an existing (unmapped) file replaces it entirely; a mapped file must not be + // rewritten, since Windows refuses to modify a file with a live mapping + Path rewritten = tempFolder.getRoot().toPath().resolve("rewritten.fvecs"); + SiftLoader.writeFvecs(rewritten, new ListRandomAccessVectorValues(vectors, 6)); + SiftLoader.writeFvecs(rewritten, new ListRandomAccessVectorValues(vectors.subList(0, 5), 6)); + assertEquals(5, new MappedFvecsRandomAccessVectorValues(rewritten).size()); + } +} diff --git a/jvector-examples/src/test/java/io/github/jbellis/jvector/example/yaml/DatasetCollectionTest.java b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/yaml/DatasetCollectionTest.java new file mode 100644 index 000000000..8414c6595 --- /dev/null +++ b/jvector-examples/src/test/java/io/github/jbellis/jvector/example/yaml/DatasetCollectionTest.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.example.yaml; + +import org.junit.Rule; +import org.junit.Test; +import org.junit.rules.TemporaryFolder; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import static org.junit.jupiter.api.Assertions.*; + +/// Tests for {@link DatasetCollection} handling of plain and structured dataset entries. +public class DatasetCollectionTest { + + @Rule + public TemporaryFolder tempFolder = new TemporaryFolder(); + + @Test + public void mixedEntriesCanonicalizeToSugaredSpecs() throws IOException { + Path file = tempFolder.newFile("datasets.yml").toPath(); + Files.writeString(file, + "small:\n" + + " - ada002-100k\n" + + " - cohere-english-v3-100k(mmap)\n" + + " - name: gecko-100k\n" + + " profile: default\n" + + " wrappers: [ mmap ]\n" + + " - name: e5-small-v2-100k\n" + + " profile: fast\n" + + " - name: e5-base-v2-100k\n" + + "empty:\n" + + "large:\n" + + " - name: cap-6M\n" + + " wrappers: mmap, memory\n" + + " - name: cohere-english-v3-10M\n" + + " wrappers:\n" + + " - mmap\n" + + " - lru: { grain: 4096, capacityMb: 512 }\n" + + " - dpr-gemma-10m(mmap,lru[grain=2048])\n"); + + var collection = DatasetCollection.load(file.toString()); + assertEquals(List.of("ada002-100k", "cohere-english-v3-100k(mmap)", "gecko-100k(mmap)", + "e5-small-v2-100k:fast", "e5-base-v2-100k"), + collection.getSection("small")); + assertTrue(collection.getSection("empty").isEmpty()); + assertEquals(List.of("cap-6M(mmap,memory)", + "cohere-english-v3-10M(mmap,lru[grain=4096,capacityMb=512])", + "dpr-gemma-10m(mmap,lru[grain=2048])"), + collection.getSection("large")); + assertEquals(8, collection.getAll().size()); + assertTrue(collection.datasetNames.containsKey("empty")); + } + + @Test + public void invalidEntriesNameTheirSection() throws IOException { + Path file = tempFolder.newFile("bad.yml").toPath(); + Files.writeString(file, + "ok:\n" + + " - ada002-100k\n" + + "broken:\n" + + " - profile: fast\n"); + var e = assertThrows(IllegalArgumentException.class, () -> DatasetCollection.load(file.toString())); + assertTrue(e.getMessage().contains("'broken'"), e.getMessage()); + } + + @Test + public void defaultCollectionStillLoads() throws IOException { + var collection = DatasetCollection.load(); + assertFalse(collection.getAll().isEmpty()); + assertTrue(collection.getSection("regression-tests").contains("cap-1M")); + } +} diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestRangeRandomAccessVectorValues.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestRangeRandomAccessVectorValues.java new file mode 100644 index 000000000..8c826c778 --- /dev/null +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/graph/TestRangeRandomAccessVectorValues.java @@ -0,0 +1,170 @@ +/* + * 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 com.carrotsearch.randomizedtesting.RandomizedTest; +import io.github.jbellis.jvector.vector.VectorizationProvider; +import io.github.jbellis.jvector.vector.types.VectorFloat; +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.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotSame; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.assertThrows; +import static org.junit.Assert.assertTrue; + +public class TestRangeRandomAccessVectorValues extends RandomizedTest { + private static final VectorTypeSupport vts = VectorizationProvider.getInstance().getVectorTypeSupport(); + private static final int DIMENSION = 4; + + private static List> randomVectors(int count) { + List> vectors = new ArrayList<>(count); + for (int i = 0; i < count; i++) { + float[] values = new float[DIMENSION]; + for (int d = 0; d < DIMENSION; d++) { + values[d] = randomFloat(); + } + vectors.add(vts.createFloatVector(values)); + } + return vectors; + } + + private static void assertVectorEquals(VectorFloat expected, VectorFloat actual) { + assertEquals(expected.length(), actual.length()); + for (int d = 0; d < expected.length(); d++) { + assertEquals(expected.get(d), actual.get(d), 0f); + } + } + + @Test + public void testRangeReadsAreRebased() { + List> vectors = randomVectors(20); + var full = new ListRandomAccessVectorValues(vectors, DIMENSION); + + var view = full.range(5, 12); + assertEquals(7, view.size()); + assertEquals(DIMENSION, view.dimension()); + assertFalse(view.isValueShared()); + for (int i = 0; i < view.size(); i++) { + assertSame(vectors.get(5 + i), view.getVector(i)); + + var dest = vts.createFloatVector(2 * DIMENSION); + view.getVectorInto(i, dest, DIMENSION); + for (int d = 0; d < DIMENSION; d++) { + assertEquals(vectors.get(5 + i).get(d), dest.get(DIMENSION + d), 0f); + } + } + } + + @Test + public void testFullAndEmptyRanges() { + List> vectors = randomVectors(10); + var full = new ListRandomAccessVectorValues(vectors, DIMENSION); + + var whole = full.range(0, full.size()); + assertEquals(full.size(), whole.size()); + for (int i = 0; i < full.size(); i++) { + assertSame(full.getVector(i), whole.getVector(i)); + } + + var empty = full.range(4, 4); + assertEquals(0, empty.size()); + assertThrows(IndexOutOfBoundsException.class, () -> empty.getVector(0)); + + var tail = full.range(full.size(), full.size()); + assertEquals(0, tail.size()); + } + + @Test + public void testRangeBoundsAreValidated() { + var full = new ListRandomAccessVectorValues(randomVectors(10), DIMENSION); + + assertThrows(IndexOutOfBoundsException.class, () -> full.range(-1, 5)); + assertThrows(IndexOutOfBoundsException.class, () -> full.range(6, 5)); + assertThrows(IndexOutOfBoundsException.class, () -> full.range(0, 11)); + assertThrows(IndexOutOfBoundsException.class, () -> full.range(11, 11)); + + var view = full.range(2, 6); + assertThrows(IndexOutOfBoundsException.class, () -> view.getVector(-1)); + assertThrows(IndexOutOfBoundsException.class, () -> view.getVector(4)); + assertThrows(IndexOutOfBoundsException.class, () -> view.getVectorInto(4, vts.createFloatVector(DIMENSION), 0)); + } + + @Test + public void testNestedRangeCollapsesToBacking() { + List> vectors = randomVectors(20); + var full = new ListRandomAccessVectorValues(vectors, DIMENSION); + + var outer = full.range(2, 12); + var inner = outer.range(3, 7); + + assertTrue(inner instanceof RangeRandomAccessVectorValues); + var range = (RangeRandomAccessVectorValues) inner; + assertEquals(5, range.fromOrdinal()); + assertEquals(9, range.toOrdinal()); + assertEquals(4, inner.size()); + for (int i = 0; i < inner.size(); i++) { + assertSame(vectors.get(5 + i), inner.getVector(i)); + } + + assertThrows(IndexOutOfBoundsException.class, () -> outer.range(0, 11)); + } + + @Test + public void testSizeIsFixedAtCreation() { + List> vectors = new ArrayList<>(randomVectors(6)); + var full = new ListRandomAccessVectorValues(vectors, DIMENSION); + + var view = full.range(2, 6); + vectors.addAll(randomVectors(4)); + + assertEquals(10, full.size()); + assertEquals(4, view.size()); + assertThrows(IndexOutOfBoundsException.class, () -> view.getVector(4)); + assertThrows(IndexOutOfBoundsException.class, () -> full.range(0, 6).range(0, 7)); + } + + @Test + public void testCopyFollowsBackingSemantics() { + List> vectors = randomVectors(8); + var unshared = new ListRandomAccessVectorValues(vectors, DIMENSION); + var unsharedView = unshared.range(1, 5); + assertSame(unsharedView, unsharedView.copy()); + + var shared = MockVectorValues.fromValues(vectors.toArray(new VectorFloat[0])); + var sharedView = shared.range(1, 5); + assertTrue(sharedView.isValueShared()); + + var copied = sharedView.copy(); + assertNotSame(sharedView, copied); + assertEquals(sharedView.size(), copied.size()); + for (int i = 0; i < sharedView.size(); i++) { + assertVectorEquals(vectors.get(1 + i), copied.getVector(i)); + } + + // a shared backing RAVV hands back one scratch reference; the view must not hide that + var first = sharedView.getVector(0); + var second = sharedView.getVector(1); + assertSame(first, second); + assertVectorEquals(vectors.get(2), second); + } +} 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..1c29d3f05 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 @@ -19,7 +19,7 @@ import io.github.jbellis.jvector.example.benchmarks.datasets.DataSet; import io.github.jbellis.jvector.example.benchmarks.datasets.DataSets; import io.github.jbellis.jvector.graph.GraphIndexBuilder; -import io.github.jbellis.jvector.graph.ListRandomAccessVectorValues; +import io.github.jbellis.jvector.graph.RandomAccessVectorValues; import org.openjdk.jmh.annotations.Benchmark; import org.openjdk.jmh.annotations.BenchmarkMode; import org.openjdk.jmh.annotations.Fork; @@ -41,13 +41,13 @@ public class GraphBuildBench { @State(Scope.Benchmark) public static class Parameters { final DataSet ds; - final ListRandomAccessVectorValues ravv; + final RandomAccessVectorValues ravv; public Parameters() { this.ds = 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()); + this.ravv = ds.getBaseRavv(); } }