From 220ec7f63197ee08d4815c0acbf90c1272c01d1f Mon Sep 17 00:00:00 2001 From: wmz7year Date: Wed, 1 May 2024 13:03:23 +0800 Subject: [PATCH] AWS Bedrock Titan embedding model add amazon.titan-embed-text-v2:0 support. --- .../titan/api/TitanEmbeddingBedrockApi.java | 6 +++- .../titan/api/TitanEmbeddingBedrockApiIT.java | 29 +++++++++++++++++-- .../embeddings/bedrock-titan-embedding.adoc | 2 +- 3 files changed, 32 insertions(+), 5 deletions(-) diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApi.java index 5901799c3..016ec4306 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApi.java @@ -153,7 +153,11 @@ public class TitanEmbeddingBedrockApi extends /** * amazon.titan-embed-text-v1 */ - TITAN_EMBED_TEXT_V1("amazon.titan-embed-text-v1"); + TITAN_EMBED_TEXT_V1("amazon.titan-embed-text-v1"), + /** + * amazon.titan-embed-text-v2 + */ + TITAN_EMBED_TEXT_V2("amazon.titan-embed-text-v2:0");; private final String id; diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApiIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApiIT.java index a666793e0..cffc00569 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApiIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/api/TitanEmbeddingBedrockApiIT.java @@ -21,6 +21,8 @@ import java.util.Base64; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; + +import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider; import software.amazon.awssdk.regions.Region; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingModel; @@ -28,20 +30,24 @@ import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEm import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingResponse; import org.springframework.core.io.DefaultResourceLoader; +import com.fasterxml.jackson.databind.ObjectMapper; + import static org.assertj.core.api.Assertions.assertThat; /** * @author Christian Tzolov + * @author Wei Jiang */ @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") public class TitanEmbeddingBedrockApiIT { @Test - public void embedText() { + public void embedTextV1() { TitanEmbeddingBedrockApi titanEmbedApi = new TitanEmbeddingBedrockApi( - TitanEmbeddingModel.TITAN_EMBED_TEXT_V1.id(), Region.US_EAST_1.id(), Duration.ofMinutes(2)); + TitanEmbeddingModel.TITAN_EMBED_TEXT_V1.id(), EnvironmentVariableCredentialsProvider.create(), + Region.US_EAST_1.id(), new ObjectMapper(), Duration.ofMinutes(2)); TitanEmbeddingRequest request = TitanEmbeddingRequest.builder().withInputText("I like to eat apples.").build(); @@ -52,11 +58,28 @@ public class TitanEmbeddingBedrockApiIT { assertThat(response.embedding()).hasSize(1536); } + @Test + public void embedTextV2() { + + TitanEmbeddingBedrockApi titanEmbedApi = new TitanEmbeddingBedrockApi( + TitanEmbeddingModel.TITAN_EMBED_TEXT_V2.id(), EnvironmentVariableCredentialsProvider.create(), + Region.US_WEST_2.id(), new ObjectMapper(), Duration.ofMinutes(2)); + + TitanEmbeddingRequest request = TitanEmbeddingRequest.builder().withInputText("I like to eat apples.").build(); + + TitanEmbeddingResponse response = titanEmbedApi.embedding(request); + + assertThat(response).isNotNull(); + assertThat(response.inputTextTokenCount()).isEqualTo(7); + assertThat(response.embedding()).hasSize(1024); + } + @Test public void embedImage() throws IOException { TitanEmbeddingBedrockApi titanEmbedApi = new TitanEmbeddingBedrockApi( - TitanEmbeddingModel.TITAN_EMBED_IMAGE_V1.id(), Region.US_EAST_1.id(), Duration.ofMinutes(2)); + TitanEmbeddingModel.TITAN_EMBED_IMAGE_V1.id(), EnvironmentVariableCredentialsProvider.create(), + Region.US_EAST_1.id(), new ObjectMapper(), Duration.ofMinutes(2)); byte[] image = new DefaultResourceLoader().getResource("classpath:/spring_framework.png") .getContentAsByteArray(); diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc index b7bc8a74e..e1a5a6600 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc @@ -78,7 +78,7 @@ The prefix `spring.ai.bedrock.titan.embedding` (defined in `BedrockTitanEmbeddin | spring.ai.bedrock.titan.embedding.model | The model id to use. See the `TitanEmbeddingModel` for the supported models. | amazon.titan-embed-image-v1 |==== -Supported values are: `amazon.titan-embed-image-v1` and `amazon.titan-embed-text-v1`. +Supported values are: `amazon.titan-embed-image-v1`, `amazon.titan-embed-text-v1` and `amazon.titan-embed-text-v2:0`. Model ID values can also be found in the https://docs.aws.amazon.com/bedrock/latest/userguide/model-ids-arns.html[AWS Bedrock documentation for base model IDs]. == Runtime Options [[embedding-options]]