Replace the Embedding format from List<Double> to float[]
- Adjust all affected classes including the Document. - Update docs. Related to #405
This commit is contained in:
committed by
Mark Pollack
parent
656fa8b4fe
commit
d538e00643
@@ -31,6 +31,7 @@ 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.model.EmbeddingUtils;
|
||||
import org.springframework.ai.model.ModelOptionsUtils;
|
||||
import org.springframework.beans.factory.InitializingBean;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
@@ -57,22 +58,23 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
|
||||
|
||||
PG_ARRAY("", null, (rs, i) -> {
|
||||
Array embedding = rs.getArray("embedding");
|
||||
return Arrays.stream((Float[]) embedding.getArray()).map(Float::doubleValue).toList();
|
||||
return EmbeddingUtils.toPrimitive((Float[]) embedding.getArray());
|
||||
|
||||
}),
|
||||
|
||||
PG_VECTOR("::vector", "vector", (rs, i) -> {
|
||||
String embedding = rs.getString("embedding");
|
||||
return Arrays.stream((embedding.substring(1, embedding.length() - 1)
|
||||
/* remove leading '[' and trailing ']' */.split(","))).map(Double::parseDouble).toList();
|
||||
return EmbeddingUtils.toPrimitive(Arrays.stream((embedding.substring(1, embedding.length() - 1)
|
||||
/* remove leading '[' and trailing ']' */.split(","))).map(Float::parseFloat).toList());
|
||||
});
|
||||
|
||||
private final String cast;
|
||||
|
||||
private final String extensionName;
|
||||
|
||||
private final RowMapper<List<Double>> rowMapper;
|
||||
private final RowMapper<float[]> rowMapper;
|
||||
|
||||
VectorType(String cast, String extensionName, RowMapper<List<Double>> rowMapper) {
|
||||
VectorType(String cast, String extensionName, RowMapper<float[]> rowMapper) {
|
||||
this.cast = cast;
|
||||
this.extensionName = extensionName;
|
||||
this.rowMapper = rowMapper;
|
||||
@@ -156,7 +158,7 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
|
||||
|
||||
@SuppressWarnings("null")
|
||||
@Override
|
||||
public List<Double> embed(String text) {
|
||||
public float[] embed(String text) {
|
||||
return this.jdbcTemplate.queryForObject(
|
||||
"SELECT pgml.embed(?, ?, ?::JSONB)" + this.defaultOptions.getVectorType().cast + " AS embedding",
|
||||
this.defaultOptions.getVectorType().rowMapper, this.defaultOptions.getTransformer(), text,
|
||||
@@ -164,7 +166,7 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<Double> embed(Document document) {
|
||||
public float[] embed(Document document) {
|
||||
return this.embed(document.getFormattedContent(this.defaultOptions.getMetadataMode()));
|
||||
}
|
||||
|
||||
@@ -175,7 +177,7 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
|
||||
final PostgresMlEmbeddingOptions optionsToUse = this.mergeOptions(request.getOptions());
|
||||
|
||||
List<Embedding> data = new ArrayList<>();
|
||||
List<List<Double>> embed = List.of();
|
||||
List<float[]> embed = List.of();
|
||||
|
||||
List<String> texts = request.getInstructions();
|
||||
if (!CollectionUtils.isEmpty(texts)) {
|
||||
@@ -187,7 +189,7 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements
|
||||
preparedStatement.setArray(3, connection.createArrayOf("TEXT", texts.toArray(Object[]::new)));
|
||||
return preparedStatement;
|
||||
}, rs -> {
|
||||
List<List<Double>> result = new ArrayList<>();
|
||||
List<float[]> result = new ArrayList<>();
|
||||
while (rs.next()) {
|
||||
result.add(optionsToUse.getVectorType().rowMapper.mapRow(rs, -1));
|
||||
}
|
||||
|
||||
@@ -84,7 +84,7 @@ class PostgresMlEmbeddingModelIT {
|
||||
PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate);
|
||||
embeddingModel.afterPropertiesSet();
|
||||
|
||||
List<Double> embed = embeddingModel.embed("Hello World!");
|
||||
float[] embed = embeddingModel.embed("Hello World!");
|
||||
|
||||
assertThat(embed).hasSize(768);
|
||||
}
|
||||
@@ -98,7 +98,7 @@ class PostgresMlEmbeddingModelIT {
|
||||
.build());
|
||||
embeddingModel.afterPropertiesSet();
|
||||
|
||||
List<Double> embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
float[] embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
|
||||
assertThat(embed).hasSize(768);
|
||||
}
|
||||
@@ -109,7 +109,7 @@ class PostgresMlEmbeddingModelIT {
|
||||
PostgresMlEmbeddingOptions.builder().withTransformer("intfloat/e5-small").build());
|
||||
embeddingModel.afterPropertiesSet();
|
||||
|
||||
List<Double> embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
float[] embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
|
||||
assertThat(embed).hasSize(384);
|
||||
}
|
||||
@@ -125,7 +125,7 @@ class PostgresMlEmbeddingModelIT {
|
||||
.build());
|
||||
embeddingModel.afterPropertiesSet();
|
||||
|
||||
List<Double> embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
float[] embed = embeddingModel.embed(new Document("Hello World!"));
|
||||
|
||||
assertThat(embed).hasSize(768);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user