initial refactoring

This commit is contained in:
Mark Pollack
2024-07-15 17:45:08 -04:00
parent ed815d8b3d
commit 53902e6fe2
39 changed files with 598 additions and 503 deletions

View File

@@ -200,9 +200,13 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
}
}
var metadata = new EmbeddingResponseMetadata(
Map.of("transformer", optionsToUse.getTransformer(), "vector-type", optionsToUse.getVectorType().name(),
"kwargs", ModelOptionsUtils.toJsonString(optionsToUse.getKwargs())));
var metadata = new EmbeddingResponseMetadata();
Map<String, Object> embeddingMetadata = Map.of("transformer", optionsToUse.getTransformer(), "vector-type",
optionsToUse.getVectorType().name(), "kwargs",
ModelOptionsUtils.toJsonString(optionsToUse.getKwargs()));
for (Map.Entry<String, Object> entry : embeddingMetadata.entrySet()) {
metadata.put(entry.getKey(), entry.getValue());
}
return new EmbeddingResponse(data, metadata);
}

View File

@@ -30,6 +30,7 @@ import org.junit.jupiter.params.provider.ValueSource;
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.ai.postgresml.PostgresMlEmbeddingModel.VectorType;
import org.testcontainers.containers.PostgreSQLContainer;
@@ -144,7 +145,21 @@ class PostgresMlEmbeddingModelIT {
assertThat(embeddingResponse).isNotNull();
assertThat(embeddingResponse.getResults()).hasSize(3);
assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf(
EmbeddingResponseMetadata metadata = embeddingResponse.getMetadata();
assertThat(metadata.get("transformer").toString())
.as("Transformer in metadata should be 'distilbert-base-uncased'")
.isEqualTo("distilbert-base-uncased");
assertThat(metadata.get("vector-type").toString())
.as("Vector type in metadata should match expected vector type")
.isEqualTo(vectorType);
assertThat(metadata.get("kwargs").toString()).as("kwargs in metadata should be '{}'").isEqualTo("{}");
assertThat(metadata.getRawMap().keySet()).as("Metadata should contain exactly the expected keys")
.containsExactlyInAnyOrder("transformer", "vector-type", "kwargs");
assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf(
Map.of("transformer", "distilbert-base-uncased", "vector-type", vectorType, "kwargs", "{}"));
assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0);
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768);
@@ -170,7 +185,7 @@ class PostgresMlEmbeddingModelIT {
assertThat(embeddingResponse).isNotNull();
assertThat(embeddingResponse.getResults()).hasSize(3);
assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer",
assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer",
"distilbert-base-uncased", "vector-type", VectorType.PG_VECTOR.name(), "kwargs", "{}"));
assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0);
assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768);
@@ -192,7 +207,7 @@ class PostgresMlEmbeddingModelIT {
assertThat(embeddingResponse).isNotNull();
assertThat(embeddingResponse.getResults()).hasSize(3);
assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer",
assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer",
"intfloat/e5-small", "vector-type", VectorType.PG_ARRAY.name(), "kwargs", "{\"device\":\"cpu\"}"));
assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0);