diff --git a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java index d86460526..f09d8a22e 100644 --- a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java +++ b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java @@ -58,9 +58,24 @@ import ai.onnxruntime.OrtSession; import io.micrometer.observation.ObservationRegistry; /** - * https://www.sbert.net/index.html https://www.sbert.net/docs/pretrained_models.html + * An implementation of the AbstractEmbeddingModel that uses ONNX-based Transformer models + * for text embeddings. + * + *

+ * By default, it uses the all-MiniLM-L6-v2 model, but can be configured to use other + * ONNX-compatible models. The class supports both CPU and GPU inference, caching of model + * resources, and various tokenization options. + *

+ * + *

+ * For more information on the underlying SBERT framework, see: + * SBERT Documentation + * SBERT Pre-trained + * Models + *

* * @author Christian Tzolov + * @since 1.0.0 */ public class TransformersEmbeddingModel extends AbstractEmbeddingModel implements InitializingBean { @@ -209,13 +224,14 @@ public class TransformersEmbeddingModel extends AbstractEmbeddingModel implement // onnxruntime this.environment = OrtEnvironment.getEnvironment(); - var sessionOptions = new OrtSession.SessionOptions(); - if (this.gpuDeviceId >= 0) { - sessionOptions.addCUDA(this.gpuDeviceId); // Run on a GPU or with another - // provider + try (var sessionOptions = new OrtSession.SessionOptions()) { + if (this.gpuDeviceId >= 0) { + sessionOptions.addCUDA(this.gpuDeviceId); // Run on a GPU or with another + // provider + } + this.session = this.environment.createSession(getCachedResource(this.modelResource).getContentAsByteArray(), + sessionOptions); } - this.session = this.environment.createSession(getCachedResource(this.modelResource).getContentAsByteArray(), - sessionOptions); this.onnxModelInputs = this.session.getInputNames(); Set onnxModelOutputs = this.session.getOutputNames(); @@ -291,39 +307,44 @@ public class TransformersEmbeddingModel extends AbstractEmbeddingModel implement token_type_ids0[i] = encodings[i].getTypeIds(); } - OnnxTensor inputIds = OnnxTensor.createTensor(this.environment, input_ids0); - OnnxTensor attentionMask = OnnxTensor.createTensor(this.environment, attention_mask0); - OnnxTensor tokenTypeIds = OnnxTensor.createTensor(this.environment, token_type_ids0); + try (OnnxTensor inputIds = OnnxTensor.createTensor(this.environment, input_ids0); + OnnxTensor attentionMask = OnnxTensor.createTensor(this.environment, attention_mask0); + OnnxTensor tokenTypeIds = OnnxTensor.createTensor(this.environment, token_type_ids0);) { - Map modelInputs = Map.of("input_ids", inputIds, "attention_mask", attentionMask, - "token_type_ids", tokenTypeIds); + Map modelInputs = Map.of("input_ids", inputIds, "attention_mask", + attentionMask, "token_type_ids", tokenTypeIds); - modelInputs = removeUnknownModelInputs(modelInputs); + modelInputs = removeUnknownModelInputs(modelInputs); - // The Run result object is AutoCloseable to prevent references from - // leaking - // out. Once the Result object is - // closed, all it’s child OnnxValues are closed too. - try (OrtSession.Result results = this.session.run(modelInputs)) { + // The Run result object is AutoCloseable to prevent references + // from leaking out. Once the Result object is + // closed, all it’s child OnnxValues are closed too. + try (OrtSession.Result results = this.session.run(modelInputs)) { - // OnnxValue lastHiddenState = results.get(0); - OnnxValue lastHiddenState = results.get(this.modelOutputName).get(); + // OnnxValue lastHiddenState = results.get(0); + OnnxValue lastHiddenState = results.get(this.modelOutputName).get(); - // 0 - batch_size (1..x) - // 1 - sequence_length (128) - // 2 - embedding dimensions (384) - float[][][] tokenEmbeddings = (float[][][]) lastHiddenState.getValue(); + // 0 - batch_size (1..x) + // 1 - sequence_length (128) + // 2 - embedding dimensions (384) + float[][][] tokenEmbeddings = (float[][][]) lastHiddenState.getValue(); - try (NDManager manager = NDManager.newBaseManager()) { - NDArray ndTokenEmbeddings = create(tokenEmbeddings, manager); - NDArray ndAttentionMask = manager.create(attention_mask0); + try (NDManager manager = NDManager.newBaseManager()) { + NDArray ndTokenEmbeddings = create(tokenEmbeddings, manager); + NDArray ndAttentionMask = manager.create(attention_mask0); - NDArray embedding = meanPooling(ndTokenEmbeddings, ndAttentionMask); + NDArray embedding = meanPooling(ndTokenEmbeddings, ndAttentionMask); - for (int i = 0; i < embedding.size(0); i++) { - resultEmbeddings.add(embedding.get(i).toFloatArray()); + for (int i = 0; i < embedding.size(0); i++) { + resultEmbeddings.add(embedding.get(i).toFloatArray()); + } } } + finally { + inputIds.close(); + attentionMask.close(); + tokenTypeIds.close(); + } } } catch (OrtException ex) {