Align the Embedding api with the new meta-model

- Craete new EmbeddingOptions -> ModelOptions, EmbeddigRequest -> ModelRequest, EmbeddingResponseMetadata -> ResponseMetadata and EmbeddignResultMetadata -> ResultMetadata.
 - Make the EmbeddigClient interface extend from ModelClient<EmbeddingRequest, EmbeddingResponse>, EmbeddingResponse implements ModelResponise and Embedding implements ModelResult.
 - Fix affected tests.
 - Steramline the EmbeddingClient interface with default method implementations based on call.
 - Merge EmbeddingUtil into AbstractEmbeddingClient
This commit is contained in:
Christian Tzolov
2024-01-25 15:13:51 +01:00
committed by Mark Pollack
parent 85ed261c08
commit 5d55b68380
32 changed files with 471 additions and 319 deletions

View File

@@ -14,7 +14,10 @@ import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingClient;
import org.springframework.ai.embedding.Embedding;
import org.springframework.ai.embedding.EmbeddingOptions;
import org.springframework.ai.embedding.EmbeddingRequest;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.embedding.EmbeddingResponseMetadata;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.core.RowMapper;
@@ -153,13 +156,20 @@ public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implement
@Override
public EmbeddingResponse embedForResponse(List<String> texts) {
return this.call(new EmbeddingRequest(texts, new EmbeddingOptions()));
}
@Override
public EmbeddingResponse call(EmbeddingRequest request) {
List<Embedding> data = new ArrayList<>();
List<List<Double>> embed = this.embed(texts);
List<List<Double>> embed = this.embed(request.getInstructions());
for (int i = 0; i < embed.size(); i++) {
data.add(new Embedding(embed.get(i), i));
}
return new EmbeddingResponse(data,
var metadata = new EmbeddingResponseMetadata(
Map.of("transformer", this.transformer, "vector-type", this.vectorType.name(), "kwargs", this.kwargs));
return new EmbeddingResponse(data, metadata);
}
@Override

View File

@@ -108,15 +108,15 @@ class PostgresMlEmbeddingClientIT {
EmbeddingResponse embeddingResponse = embeddingClient
.embedForResponse(List.of("Hello World!", "Spring AI!", "LLM!"));
assertThat(embeddingResponse).isNotNull();
assertThat(embeddingResponse.getData()).hasSize(3);
assertThat(embeddingResponse.getResults()).hasSize(3);
assertThat(embeddingResponse.getMetadata()).containsExactlyEntriesOf(
Map.of("transformer", "distilbert-base-uncased", "vector-type", vectorType, "kwargs", "{}"));
assertThat(embeddingResponse.getData().get(0).getIndex()).isEqualTo(0);
assertThat(embeddingResponse.getData().get(0).getEmbedding()).hasSize(768);
assertThat(embeddingResponse.getData().get(1).getIndex()).isEqualTo(1);
assertThat(embeddingResponse.getData().get(1).getEmbedding()).hasSize(768);
assertThat(embeddingResponse.getData().get(2).getIndex()).isEqualTo(2);
assertThat(embeddingResponse.getData().get(2).getEmbedding()).hasSize(768);
assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0);
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768);
assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1);
assertThat(embeddingResponse.getResults().get(1).getOutput()).hasSize(768);
assertThat(embeddingResponse.getResults().get(2).getIndex()).isEqualTo(2);
assertThat(embeddingResponse.getResults().get(2).getOutput()).hasSize(768);
// embeddingClient.dropPgmlExtension();
}