From 65a6ffcffb11ce9cdfb83f021f141c40a26fbf03 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Thu, 28 Sep 2023 08:57:09 +0200 Subject: [PATCH] Dynamic embedding dimensions resolution Leverage #28 to allow the vector stores to resolve the embedding dimensions dynamically. The explicitly set dimensions (if set) precedence over other configurations. If the embedding dimensions are not explicitly set, the embeddingClient is used to determine them dynamically. If the client fails the it falls back to 1536. --- .../embedding/AzureOpenAiEmbeddingClient.java | 14 +-- ...gUtilTest.java => EmbeddingUtilTests.java} | 2 +- .../ai/vectorstore/MilvusVectorStore.java | 23 ++++- .../MilvusEmbeddingDimensionsTests.java | 87 +++++++++++++++++++ .../ai/vectorstore/MilvusVectorStoreIT.java | 9 +- .../ai/vectorstore/PgVectorStore.java | 36 +++++++- .../PgVectorEmbeddingDimensionsTests.java | 78 +++++++++++++++++ .../ai/vectorstore/PgVectorStoreIT.java | 12 +-- 8 files changed, 237 insertions(+), 24 deletions(-) rename spring-ai-core/src/test/java/org/springframework/ai/embedding/{EmbeddingUtilTest.java => EmbeddingUtilTests.java} (98%) create mode 100644 vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java create mode 100644 vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java diff --git a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java index 79622d2bb..f84f133d6 100644 --- a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java +++ b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java @@ -1,5 +1,11 @@ package org.springframework.ai.azure.openai.embedding; +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; + import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.models.EmbeddingItem; import com.azure.ai.openai.models.Embeddings; @@ -7,6 +13,7 @@ import com.azure.ai.openai.models.EmbeddingsOptions; import com.azure.ai.openai.models.EmbeddingsUsage; import org.slf4j.Logger; import org.slf4j.LoggerFactory; + import org.springframework.ai.document.Document; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingClient; @@ -14,13 +21,6 @@ import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.embedding.EmbeddingUtil; import org.springframework.util.Assert; -import java.util.ArrayList; -import java.util.HashMap; -import java.util.List; -import java.util.Map; -import java.util.concurrent.atomic.AtomicInteger; -import java.util.stream.Collectors; - public class AzureOpenAiEmbeddingClient implements EmbeddingClient { private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiEmbeddingClient.class); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTest.java b/spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTests.java similarity index 98% rename from spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTest.java rename to spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTests.java index 7a26f2851..92cea53ea 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTest.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/embedding/EmbeddingUtilTests.java @@ -38,7 +38,7 @@ import static org.mockito.Mockito.when; * @author Christian Tzolov */ @ExtendWith(MockitoExtension.class) -public class EmbeddingUtilTest { +public class EmbeddingUtilTests { @Mock private EmbeddingClient embeddingClient; diff --git a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java index f6466346b..bcc4bebe0 100644 --- a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java +++ b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java @@ -67,6 +67,8 @@ public class MilvusVectorStore implements VectorStore, SmartLifecycle { public static final int OPENAI_EMBEDDING_DIMENSION_SIZE = 1536; + public static final int INVALID_EMBEDDING_DIMENSION = -1; + public static final String DEFAULT_DATABASE_NAME = "default"; public static final String DEFAULT_COLLECTION_NAME = "vector_store"; @@ -140,7 +142,7 @@ public class MilvusVectorStore implements VectorStore, SmartLifecycle { private String collectionName = DEFAULT_COLLECTION_NAME; - private int embeddingDimension = OPENAI_EMBEDDING_DIMENSION_SIZE; + private int embeddingDimension = INVALID_EMBEDDING_DIMENSION; private IndexType indexType = IndexType.IVF_FLAT; @@ -442,7 +444,7 @@ public class MilvusVectorStore implements VectorStore, SmartLifecycle { FieldType embeddingFieldType = FieldType.newBuilder() .withName(EMBEDDING_FIELD_NAME) .withDataType(DataType.FloatVector) - .withDimension(this.config.embeddingDimension) + .withDimension(this.embeddingDimensions()) .build(); CreateCollectionParam createCollectionReq = CreateCollectionParam.newBuilder() @@ -495,6 +497,23 @@ public class MilvusVectorStore implements VectorStore, SmartLifecycle { } } + int embeddingDimensions() { + if (this.config.embeddingDimension != INVALID_EMBEDDING_DIMENSION) { + return this.config.embeddingDimension; + } + try { + int embeddingDimensions = this.embeddingClient.dimensions(); + if (embeddingDimensions > 0) { + return embeddingDimensions; + } + } + catch (Exception e) { + logger.warn("Failed to obtain the embedding dimensions from the embedding client and fall backs to default:" + + this.config.embeddingDimension, e); + } + return OPENAI_EMBEDDING_DIMENSION_SIZE; + } + // used by the test as well void dropCollection() { diff --git a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java new file mode 100644 index 000000000..1f4bd158a --- /dev/null +++ b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java @@ -0,0 +1,87 @@ +/* + * Copyright 2023-2023 the original author or authors. + * + * 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 + * + * https://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 org.springframework.ai.vectorstore; + +import io.milvus.client.MilvusServiceClient; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.only; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * @author Christian Tzolov + */ +@ExtendWith(MockitoExtension.class) +public class MilvusEmbeddingDimensionsTests { + + @Mock + private EmbeddingClient embeddingClient; + + @Mock + private MilvusServiceClient milvusClient; + + @Test + public void explicitlySetDimensions() { + + final int explicitDimensions = 696; + + MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder() + .withEmbeddingDimension(explicitDimensions) + .build(); + + var dim = new MilvusVectorStore(milvusClient, embeddingClient, config).embeddingDimensions(); + + assertThat(dim).isEqualTo(explicitDimensions); + verify(embeddingClient, never()).dimensions(); + } + + @Test + public void embeddingClientDimensions() { + when(embeddingClient.dimensions()).thenReturn(969); + + MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder().build(); + + var dim = new MilvusVectorStore(milvusClient, embeddingClient, config).embeddingDimensions(); + + assertThat(dim).isEqualTo(969); + + verify(embeddingClient, only()).dimensions(); + } + + @Test + public void fallBackToDefaultDimensions() { + + when(embeddingClient.dimensions()).thenThrow(new RuntimeException()); + + var dim = new MilvusVectorStore(milvusClient, embeddingClient, + MilvusVectorStoreConfig.builder().build()) + .embeddingDimensions(); + + assertThat(dim).isEqualTo(MilvusVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE); + verify(embeddingClient, only()).dimensions(); + } + +} diff --git a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java index 763b296fb..185035248 100644 --- a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java +++ b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java @@ -102,7 +102,7 @@ public class MilvusVectorStoreIT { public void addAndSearchTest(String metricType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.milvus.metricType=" + metricType) + .withPropertyValues("test.spring.ai.vectorstore.milvus.metricType=" + metricType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -135,7 +135,7 @@ public class MilvusVectorStoreIT { public void documentUpdateTest(String metricType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.milvus.metricType=" + metricType) + .withPropertyValues("test.spring.ai.vectorstore.milvus.metricType=" + metricType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -181,7 +181,7 @@ public class MilvusVectorStoreIT { public void searchThresholdTest(String metricType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.milvus.metricType=" + metricType) + .withPropertyValues("test.spring.ai.vectorstore.milvus.metricType=" + metricType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -217,7 +217,7 @@ public class MilvusVectorStoreIT { @EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class }) public static class TestApplication { - @Value("${spring.ai.vectorstore.milvus.metricType}") + @Value("${test.spring.ai.vectorstore.milvus.metricType}") private MetricType metricType; @Bean @@ -227,7 +227,6 @@ public class MilvusVectorStoreIT { .withDatabaseName("default") .withIndexType(IndexType.IVF_FLAT) .withMetricType(metricType) - .withEmbeddingDimension(MilvusVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE) .build(); return new MilvusVectorStore(milvusClient, embeddingClient, config); } diff --git a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java index 6972babde..8fde75009 100644 --- a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java +++ b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java @@ -29,6 +29,8 @@ import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.ObjectMapper; import com.pgvector.PGvector; import org.postgresql.util.PGobject; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; @@ -45,8 +47,12 @@ import org.springframework.lang.Nullable; */ public class PgVectorStore implements VectorStore, SmartLifecycle { + private static final Logger logger = LoggerFactory.getLogger(PgVectorStore.class); + public static final int OPENAI_EMBEDDING_DIMENSION_SIZE = 1536; + public static final int INVALID_EMBEDDING_DIMENSION = -1; + public static final String VECTOR_TABLE_NAME = "vector_store"; private final JdbcTemplate jdbcTemplate; @@ -176,8 +182,13 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { } public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) { - this(jdbcTemplate, embeddingClient, OPENAI_EMBEDDING_DIMENSION_SIZE, - PgVectorStore.PgDistanceType.CosineDistance, false, PgIndexType.NONE); + this(jdbcTemplate, embeddingClient, INVALID_EMBEDDING_DIMENSION, PgVectorStore.PgDistanceType.CosineDistance, + false, PgIndexType.NONE); + } + + public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions) { + this(jdbcTemplate, embeddingClient, dimensions, PgVectorStore.PgDistanceType.CosineDistance, false, + PgIndexType.NONE); } public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions, @@ -299,7 +310,7 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { this.jdbcTemplate.execute("CREATE TABLE IF NOT EXISTS " + VECTOR_TABLE_NAME + " ( " + "id uuid DEFAULT uuid_generate_v4 () PRIMARY KEY, " + "content text, " + "metadata json, " - + "embedding vector(" + this.dimensions + "))"); + + "embedding vector(" + this.embeddingDimensions() + "))"); if (this.createIndexMethod != PgIndexType.NONE) { this.jdbcTemplate.execute("CREATE INDEX ON " + VECTOR_TABLE_NAME + " USING " + this.createIndexMethod @@ -309,6 +320,25 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { this.isRunning.set(true); } + int embeddingDimensions() { + // The manually set dimensions have precedence over the computed one. + if (this.dimensions > 0) { + return this.dimensions; + } + + try { + int embeddingDimensions = this.embeddingClient.dimensions(); + if (embeddingDimensions > 0) { + return embeddingDimensions; + } + } + catch (Exception e) { + logger.warn("Failed to obtain the embedding dimensions from the embedding client and fall backs to default:" + + OPENAI_EMBEDDING_DIMENSION_SIZE, e); + } + return OPENAI_EMBEDDING_DIMENSION_SIZE; + } + @Override public void stop() { // Remove existing VectorStoreTable diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java new file mode 100644 index 000000000..dc18a6997 --- /dev/null +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java @@ -0,0 +1,78 @@ +/* + * Copyright 2023-2023 the original author or authors. + * + * 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 + * + * https://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 org.springframework.ai.vectorstore; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.jdbc.core.JdbcTemplate; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.only; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * @author Christian Tzolov + */ +@ExtendWith(MockitoExtension.class) +public class PgVectorEmbeddingDimensionsTests { + + @Mock + private EmbeddingClient embeddingClient; + + @Mock + private JdbcTemplate jdbcTemplate; + + @Test + public void explicitlySetDimensions() { + + final int explicitDimensions = 696; + + var dim = new PgVectorStore(jdbcTemplate, embeddingClient, explicitDimensions).embeddingDimensions(); + + assertThat(dim).isEqualTo(explicitDimensions); + verify(embeddingClient, never()).dimensions(); + } + + @Test + public void embeddingClientDimensions() { + when(embeddingClient.dimensions()).thenReturn(969); + + var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions(); + + assertThat(dim).isEqualTo(969); + + verify(embeddingClient, only()).dimensions(); + } + + @Test + public void fallBackToDefaultDimensions() { + + when(embeddingClient.dimensions()).thenThrow(new RuntimeException()); + + var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions(); + + assertThat(dim).isEqualTo(PgVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE); + verify(embeddingClient, only()).dimensions(); + } + +} diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java index aab80b211..1a91d224e 100644 --- a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java @@ -72,7 +72,7 @@ public class PgVectorStoreIT { private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() .withUserConfiguration(TestApplication.class) .withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"), - "spring.ai.vectorstore.pgvector.distanceType=CosineDistance", + "test.spring.ai.vectorstore.pgvector.distanceType=CosineDistance", // JdbcTemplate configuration String.format("app.datasource.url=jdbc:postgresql://localhost:%d/%s", @@ -84,7 +84,7 @@ public class PgVectorStoreIT { @ValueSource(strings = { "CosineDistance", "EuclideanDistance", "NegativeInnerProduct" }) public void addAndSearchTest(String distanceType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType) + .withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + distanceType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -114,7 +114,7 @@ public class PgVectorStoreIT { public void documentUpdateTest(String distanceType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType) + .withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + distanceType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -153,7 +153,7 @@ public class PgVectorStoreIT { public void searchThresholdTest(String distanceType) { contextRunner.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class)) - .withPropertyValues("spring.ai.vectorstore.pgvector.distanceType=" + distanceType) + .withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + distanceType) .run(context -> { VectorStore vectorStore = context.getBean(VectorStore.class); @@ -210,12 +210,12 @@ public class PgVectorStoreIT { @EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class }) public static class TestApplication { - @Value("${spring.ai.vectorstore.pgvector.distanceType}") + @Value("${test.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, + return new PgVectorStore(jdbcTemplate, embeddingClient, PgVectorStore.INVALID_EMBEDDING_DIMENSION, distanceType, true, PgIndexType.HNSW); }