Set the default OpenAI Embedding model to text-embedding-3-small

The text-embedding-3-small has the same dimensions as previous text-embedding-ada-002.
  The text-embedding-3-large has higher dimensionality not supported by some Vector Stores.
This commit is contained in:
Christian Tzolov
2024-02-16 10:27:13 +01:00
parent e62543c761
commit 7ef49b8cef
6 changed files with 52 additions and 15 deletions

View File

@@ -50,8 +50,6 @@ public class OpenAiEmbeddingClient extends AbstractEmbeddingClient {
private static final Logger logger = LoggerFactory.getLogger(OpenAiEmbeddingClient.class);
public static final String DEFAULT_OPENAI_EMBEDDING_MODEL = "text-embedding-3-large";
private final OpenAiEmbeddingOptions defaultOptions;
private final RetryTemplate retryTemplate = RetryTemplate.builder()
@@ -76,7 +74,7 @@ public class OpenAiEmbeddingClient extends AbstractEmbeddingClient {
public OpenAiEmbeddingClient(OpenAiApi openAiApi, MetadataMode metadataMode) {
this(openAiApi, metadataMode,
OpenAiEmbeddingOptions.builder().withModel(DEFAULT_OPENAI_EMBEDDING_MODEL).build());
OpenAiEmbeddingOptions.builder().withModel(OpenAiApi.DEFAULT_EMBEDDING_MODEL).build());
}
public OpenAiEmbeddingClient(OpenAiApi openAiApi, MetadataMode metadataMode, OpenAiEmbeddingOptions options) {
@@ -106,7 +104,7 @@ public class OpenAiEmbeddingClient extends AbstractEmbeddingClient {
this.defaultOptions.getModel(), this.defaultOptions.getEncodingFormat(),
this.defaultOptions.getUser())
: new org.springframework.ai.openai.api.OpenAiApi.EmbeddingRequest<>(request.getInstructions(),
DEFAULT_OPENAI_EMBEDDING_MODEL);
OpenAiApi.DEFAULT_EMBEDDING_MODEL);
if (request.getOptions() != null && !EmbeddingOptions.EMPTY.equals(request.getOptions())) {
apiRequest = ModelOptionsUtils.merge(request.getOptions(), apiRequest,

View File

@@ -54,7 +54,7 @@ public class OpenAiApi {
private static final String DEFAULT_BASE_URL = "https://api.openai.com";
public static final String DEFAULT_CHAT_MODEL = "gpt-3.5-turbo";
public static final String DEFAULT_EMBEDDING_MODEL = "text-embedding-ada-002";
public static final String DEFAULT_EMBEDDING_MODEL = "text-embedding-3-small";
private static final Predicate<String> SSE_DONE_PREDICATE = "[DONE]"::equals;
private final RestClient restClient;

View File

@@ -16,8 +16,11 @@
package org.springframework.ai.openai.embedding;
import org.junit.jupiter.api.Test;
import org.springframework.ai.embedding.EmbeddingRequest;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.openai.OpenAiEmbeddingClient;
import org.springframework.ai.openai.OpenAiEmbeddingOptions;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
@@ -32,17 +35,49 @@ class EmbeddingIT {
private OpenAiEmbeddingClient embeddingClient;
@Test
void simpleEmbedding() {
void defaultEmbedding() {
assertThat(embeddingClient).isNotNull();
EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World"));
assertThat(embeddingResponse.getResults()).hasSize(1);
assertThat(embeddingResponse.getResults().get(0)).isNotNull();
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1536);
assertThat(embeddingResponse.getMetadata()).containsEntry("model", "text-embedding-ada-002");
assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2);
assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2);
assertThat(embeddingClient.dimensions()).isEqualTo(1536);
}
@Test
void embedding3Large() {
EmbeddingResponse embeddingResponse = embeddingClient.call(new EmbeddingRequest(List.of("Hello World"),
OpenAiEmbeddingOptions.builder().withModel("text-embedding-3-large").build()));
assertThat(embeddingResponse.getResults()).hasSize(1);
assertThat(embeddingResponse.getResults().get(0)).isNotNull();
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(3072);
assertThat(embeddingResponse.getMetadata()).containsEntry("model", "text-embedding-3-large");
assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2);
assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2);
assertThat(embeddingClient.dimensions()).isEqualTo(3072);
// assertThat(embeddingClient.dimensions()).isEqualTo(3072);
}
@Test
void embedding3Small() {
EmbeddingResponse embeddingResponse = embeddingClient.call(new EmbeddingRequest(List.of("Hello World"),
OpenAiEmbeddingOptions.builder().withModel("text-embedding-3-small").build()));
assertThat(embeddingResponse.getResults()).hasSize(1);
assertThat(embeddingResponse.getResults().get(0)).isNotNull();
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1536);
assertThat(embeddingResponse.getMetadata()).containsEntry("model", "text-embedding-3-small");
assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2);
assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2);
// assertThat(embeddingClient.dimensions()).isEqualTo(3072);
}
}

View File

@@ -62,7 +62,7 @@ The prefix `spring.ai.openai.embedding` is property prefix that configures the `
| spring.ai.openai.embedding.base-url | Optional overrides the spring.ai.openai.base-url to provide embedding specific url | -
| spring.ai.openai.embedding.api-key | Optional overrides the spring.ai.openai.api-key to provide embedding specific api-key | -
| spring.ai.openai.embedding.metadata-mode | Document content extraction mode. | EMBED
| spring.ai.openai.embedding.options.model | The model to use | text-embedding-3-large (other options: text-embedding-3-small, text-embedding-ada-002)
| spring.ai.openai.embedding.options.model | The model to use | text-embedding-3-small (other options: text-embedding-3-large, text-embedding-ada-002)
| spring.ai.openai.embedding.options.encodingFormat | The format to return the embeddings in. Can be either float or base64. | -
| spring.ai.openai.embedding.options.user | A unique identifier representing your end-user, which can help OpenAI to monitor and detect abuse. | -
|====

View File

@@ -39,8 +39,10 @@ import org.testcontainers.containers.wait.strategy.Wait;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.openai.OpenAiEmbeddingClient;
import org.springframework.ai.openai.OpenAiEmbeddingOptions;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig;
import org.springframework.beans.factory.annotation.Value;
@@ -59,9 +61,9 @@ import static org.assertj.core.api.Assertions.assertThat;
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
// @EnabledIfEnvironmentVariable(named = "PALM_API_KEY", matches = ".+")
public class MilvusVectorStoreIT {
@SuppressWarnings("rawtypes")
private static DockerComposeContainer milvusContainer;
private static final File TEMP_FOLDER = new File("target/test-" + UUID.randomUUID().toString());
@@ -84,6 +86,7 @@ public class MilvusVectorStoreIT {
}
}
@SuppressWarnings({ "rawtypes", "resource" })
@BeforeAll
public static void beforeAll() {
FileSystemUtils.deleteRecursively(TEMP_FOLDER);
@@ -307,9 +310,8 @@ public class MilvusVectorStoreIT {
@Bean
public EmbeddingClient embeddingClient() {
// return new VertexAiEmbeddingClient(new
// VertexAiApi(System.getenv("PALM_API_KEY")));
return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY")), MetadataMode.EMBED,
OpenAiEmbeddingOptions.builder().withModel("text-embedding-ada-002").build());
}
}

View File

@@ -20,7 +20,7 @@ services:
retries: 3
minio:
image: minio/minio:RELEASE.2023-11-11T08-14-41Z
image: minio/minio:RELEASE.2023-03-20T20-16-18Z
environment:
MINIO_ACCESS_KEY: minioadmin
MINIO_SECRET_KEY: minioadmin
@@ -37,8 +37,10 @@ services:
retries: 3
standalone:
image: milvusdb/milvus:v2.3.5
image: milvusdb/milvus:v2.3.8
command: ["milvus", "run", "standalone"]
security_opt:
- seccomp:unconfined
environment:
ETCD_ENDPOINTS: etcd:2379
MINIO_ADDRESS: minio:9000
@@ -59,4 +61,4 @@ services:
networks:
default:
name: milvus
name: milvus