PgVectorStore similarity search threshold fix

Resolves #30
This commit is contained in:
Christian Tzolov
2023-09-22 18:54:45 +02:00
parent c0aaa63411
commit 9beb7c6b7b
2 changed files with 102 additions and 77 deletions

View File

@@ -95,19 +95,25 @@ public class PgVectorStore implements VectorStore, SmartLifecycle {
public enum PgDistanceType {
EuclideanDistance("<->", "vector_l2_ops"),
EuclideanDistance("<->", "vector_l2_ops",
"SELECT *, embedding <-> ? AS distance FROM %s WHERE embedding <-> ? < ? ORDER BY distance LIMIT ? "),
NegativeInnerProduct("<#>", "vector_ip_ops"),
NegativeInnerProduct("<#>", "vector_ip_ops",
"SELECT *, (1 + (embedding <#> ?)) AS distance FROM %s WHERE (1 + (embedding <#> ?)) < ? ORDER BY distance LIMIT ? "),
CosineDistance("<=>", "vector_cosine_ops");
CosineDistance("<=>", "vector_cosine_ops",
"SELECT *, embedding <=> ? AS distance FROM %s WHERE embedding <=> ? < ? ORDER BY distance LIMIT ? ");
public final String operator;
public final String index;
PgDistanceType(String operator, String index) {
public final String similaritySearchSqlTemplate;
PgDistanceType(String operator, String index, String sqlTemplate) {
this.operator = operator;
this.index = index;
this.similaritySearchSqlTemplate = sqlTemplate;
}
}
@@ -197,7 +203,7 @@ public class PgVectorStore implements VectorStore, SmartLifecycle {
document.setEmbedding(embedding);
UUID id = UUID.fromString(document.getId());
String content = document.getText(); // TODO: text of text + metadata?
String content = document.getText();
Map<String, Object> metadata = document.getMetadata();
PGvector pgEmbedding = new PGvector(toFloatArray(embedding));
@@ -235,20 +241,21 @@ public class PgVectorStore implements VectorStore, SmartLifecycle {
}
@Override
public List<Document> similaritySearch(String query, int k) {
PGvector queryEmbedding = getQueryEmbedding(query);
return this.jdbcTemplate.query("SELECT *, embedding " + this.comparisonOperator() + " ? AS distance FROM "
+ VECTOR_TABLE_NAME + " ORDER BY distance LIMIT ?", new DocumentRowMapper(this.objectMapper),
queryEmbedding, k);
public List<Document> similaritySearch(String query, int topK) {
return this.similaritySearch(query, topK, 0.0 /** ALL */
);
}
@Override
public List<Document> similaritySearch(String query, int k, double threshold) {
public List<Document> similaritySearch(String query, int topK, double similarityThreshold) {
double distance = 1 - similarityThreshold;
PGvector queryEmbedding = getQueryEmbedding(query);
return this.jdbcTemplate.query(
"SELECT *, embedding " + this.comparisonOperator() + " ? AS distance FROM " + VECTOR_TABLE_NAME
+ " WHERE embedding " + this.comparisonOperator() + " ? < ? ORDER BY distance LIMIT ? ",
new DocumentRowMapper(this.objectMapper), queryEmbedding, queryEmbedding, threshold, k);
String.format(this.getDistanceType().similaritySearchSqlTemplate, VECTOR_TABLE_NAME),
new DocumentRowMapper(this.objectMapper), queryEmbedding, queryEmbedding, distance, topK);
}
public List<Double> embeddingDistance(String query) {

View File

@@ -25,7 +25,8 @@ import java.util.stream.Collectors;
import javax.sql.DataSource;
import com.zaxxer.hikari.HikariDataSource;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.ValueSource;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
@@ -34,6 +35,7 @@ import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.vectorstore.PgVectorStore.PgIndexType;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
@@ -70,8 +72,8 @@ public class PgVectorStoreIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(TestApplication.class)
.withPropertyValues("spring.datasource.type=com.zaxxer.hikari.HikariDataSource",
"spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"),
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"),
"spring.ai.vectorstore.pgvector.distanceType=CosineDistance",
// JdbcTemplate configuration
String.format("app.datasource.url=jdbc:postgresql://localhost:%d/%s",
@@ -79,97 +81,110 @@ public class PgVectorStoreIT {
"app.datasource.username=postgres", "app.datasource.password=postgres",
"app.datasource.type=com.zaxxer.hikari.HikariDataSource");
@Test
public void addAndSearchTest() {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)).run(context -> {
@ParameterizedTest
@ValueSource(strings = { "CosineDistance", "EuclideanDistance", "NegativeInnerProduct" })
public void addAndSearchTest(String distanceType) {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType)
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.add(documents);
vectorStore.add(documents);
List<Document> results = vectorStore.similaritySearch("Great", 1);
List<Document> results = vectorStore.similaritySearch("Great", 1);
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getText()).isEqualTo(
"Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getText()).isEqualTo(
"Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).collect(Collectors.toList()));
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).collect(Collectors.toList()));
List<Document> results2 = vectorStore.similaritySearch("Great", 1);
assertThat(results2).hasSize(0);
List<Document> results2 = vectorStore.similaritySearch("Great", 1);
assertThat(results2).hasSize(0);
});
});
}
@Test
public void documentUpdateTest() {
@ParameterizedTest
@ValueSource(strings = { "CosineDistance", "EuclideanDistance", "NegativeInnerProduct" })
public void documentUpdateTest(String distanceType) {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)).run(context -> {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType)
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
VectorStore vectorStore = context.getBean(VectorStore.class);
Document document = new Document(UUID.randomUUID().toString(), "Spring AI rocks!!",
Collections.singletonMap("meta1", "meta1"));
Document document = new Document(UUID.randomUUID().toString(), "Spring AI rocks!!",
Collections.singletonMap("meta1", "meta1"));
vectorStore.add(List.of(document));
vectorStore.add(List.of(document));
List<Document> results = vectorStore.similaritySearch("Spring", 5);
List<Document> results = vectorStore.similaritySearch("Spring", 5);
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(document.getId());
assertThat(resultDoc.getText()).isEqualTo("Spring AI rocks!!");
assertThat(resultDoc.getMetadata()).containsKeys("meta1", "distance");
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(document.getId());
assertThat(resultDoc.getText()).isEqualTo("Spring AI rocks!!");
assertThat(resultDoc.getMetadata()).containsKeys("meta1", "distance");
Document sameIdDocument = new Document(document.getId(),
"The World is Big and Salvation Lurks Around the Corner",
Collections.singletonMap("meta2", "meta2"));
Document sameIdDocument = new Document(document.getId(),
"The World is Big and Salvation Lurks Around the Corner",
Collections.singletonMap("meta2", "meta2"));
vectorStore.add(List.of(sameIdDocument));
vectorStore.add(List.of(sameIdDocument));
results = vectorStore.similaritySearch("FooBar", 5);
results = vectorStore.similaritySearch("FooBar", 5);
assertThat(results).hasSize(1);
resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(document.getId());
assertThat(resultDoc.getText()).isEqualTo("The World is Big and Salvation Lurks Around the Corner");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
});
assertThat(results).hasSize(1);
resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(document.getId());
assertThat(resultDoc.getText()).isEqualTo("The World is Big and Salvation Lurks Around the Corner");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
});
}
@Test
public void searchThresholdTest() {
@ParameterizedTest
@ValueSource(strings = { "CosineDistance", "EuclideanDistance", "NegativeInnerProduct" })
public void searchThresholdTest(String distanceType) {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)).run(context -> {
contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType)
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.add(documents);
vectorStore.add(documents);
List<Document> fullResult = vectorStore.similaritySearch("Great", 5, 1.0);
List<Document> fullResult = vectorStore.similaritySearch("Great", 5, 0.0);
assertThat(fullResult).hasSize(3);
List<Float> distances = fullResult.stream()
.map(doc -> (Float) doc.getMetadata().get("distance"))
.collect(Collectors.toList());
assertThat(isSortedByDistance(fullResult)).isTrue();
assertThat(fullResult).hasSize(3);
fullResult.stream().forEach(doc -> System.out.println(doc.getMetadata().get("distance")));
assertThat(isSortedByDistance(fullResult)).isTrue();
List<Double> embeddingDistance = ((PgVectorStore) vectorStore).embeddingDistance("Great");
fullResult.stream().forEach(doc -> System.out.println(doc.getMetadata().get("distance")));
List<Document> results = vectorStore.similaritySearch("Great", 5, 0.21);
List<Double> embeddingDistance = ((PgVectorStore) vectorStore).embeddingDistance("Great");
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getText()).isEqualTo(
"Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
List<Document> results = vectorStore.similaritySearch("Great", 5, (1 - (distances.get(0) + 0.01)));
});
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getText()).isEqualTo(
"Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression");
assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance");
});
}
private static boolean isSortedByDistance(List<Document> docs) {
@@ -198,10 +213,13 @@ public class PgVectorStoreIT {
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class TestApplication {
@Value("${spring.ai.vectorstore.pgvector.distanceType}")
PgVectorStore.PgDistanceType distanceType;
@Bean
public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) {
return new PgVectorStore(jdbcTemplate, embeddingClient, PgVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE,
PgVectorStore.PgDistanceType.CosineDistance, true, PgIndexType.HNSW);
distanceType, true, PgIndexType.HNSW);
}
@Bean