在Java代码中使用Embedding模型实践(不依赖Python运行环境)
前言
背景:很多开源embedding模型是使用python技术栈的,通常情况下需要使用pyhon来开发接口供外部使用,但是服务通信是有开销的,加上很多业务代码都是Java代码,技术栈也是Java,为了让Java服务能用上embedding模型从而去学习和开发一套pythone服务是不划算的。本教程提供了一个将python才能用的模型迁移到Java中使用的教程,以BAAI/bge-small-zh-v1.5模型为例。
模型介绍:BAAI/bge-small-zh-v1.5是BGE开源的轻量中文语义模型,占用非常小,主要用来对比文本相似度,生成语义向量等。这个用途非常广,比如可以用这个模型生成向量落库到ES做语义搜索。更重要的它的开源协议是MIT License,这意味着只要写上来源,是可以免费商用的,并且不需要你公开源代码。详细参见:FlagEmbedding
准备
-
虽然最终使用模型不依赖python技术栈,但是模型需要从python脚本中导出,所以需要准备一个python3的环境。
-
为了方便导出,需要一个linux系统环境(虚拟机也可以),本教程用的Ubuntu24
下载模型
FlagEmbedding/README_zh.md 根据教程我们我们先把模型安装到本地
#拉取源码并安装模型依赖到虚拟环境中
git clone https://github.com/FlagOpen/FlagEmbedding.git
#创建虚拟环境并使用防止影响系统依赖
python3 -m venv .venv
source .venv/bin/activate
#安装
cd FlagEmbedding
pip install transformers==4.31.0
pip install .
#使用低版本的transformers,高版本有个函数缺失
pip install "transformers>=4.44.2,<5.0"
安装过程会比较久,过程中可以去干其他事情。安装完成之后切换到上级目录,并新建一个文件夹。然后新建一个用于导出模型的脚本。
cd ..
mkdir exoprt_model
cd exoprt_model
vim save_origin_model.py
脚本代码如下:
from FlagEmbedding import FlagAutoModel
model = FlagAutoModel.from_finetuned('BAAI/bge-small-zh',
query_instruction_for_retrieval="Represent this sentence for searching relevant passages:",
use_fp16=True)
# 2. 保存为标准 HF 格式(包含 config.json, pytorch_model.bin, tokenizer 等)
output_dir = "./bge-small-zh-v1.5-local" # 本地保存路径
model.model.save_pretrained(output_dir) # 保存 PyTorch 模型
model.tokenizer.save_pretrained(output_dir) # 保存 tokenizer
print(f"模型已保存到: {output_dir}")
保存脚本后执行如下指令:
python3 save_origin_model.py
如果报错可以尝试删缓存
rm -rf ~/.cache/huggingface/*
然后目录下会多出一个bge-small-zh-v1.5-local文件夹,这个是HF格式模型数据。我们要把他转换成ONNX格式的
将模型转换成ONNX格式
转换之前,我们需要安装一些工具
pip install optimum[onnxruntime] onnx onnxruntime sentence-transformers torch
安装完工具之后,执行如下指令将模型导出为ONNX格式
optimum-cli export onnx \
--model ./bge-small-zh-v1.5-local \
--task feature-extraction \
./onnx-bge-small
时间比较长,等待指令执行完成之后就导出成功了,成功后会看到目录下有一个onnx-bge-small文件夹,这就是我们需要的ONNX格式模型。
在Java代码中使用ONNX模型
我们新建一个java项目,并导入下面这个pom文件中的依赖,这里我用的java21,jdk的版本应该随意就行。
<!-- pom.xml -->
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<groupId>com.mocrash</groupId>
<artifactId>java-embedding-onnx</artifactId>
<version>1.0-SNAPSHOT</version>
<dependencies>
<!-- DJL 核心 API -->
<dependency>
<groupId>ai.djl</groupId>
<artifactId>api</artifactId>
<version>0.34.0</version>
</dependency>
<!-- ONNX Runtime 引擎 -->
<dependency>
<groupId>ai.djl.onnxruntime</groupId>
<artifactId>onnxruntime-engine</artifactId>
<version>0.34.0</version>
</dependency>
<!-- HuggingFace Tokenizer 支持 -->
<dependency>
<groupId>ai.djl.huggingface</groupId>
<artifactId>tokenizers</artifactId>
<version>0.34.0</version>
</dependency>
<!--<dependency>
<groupId>org.slf4j</groupId>
<artifactId>slf4j-simple</artifactId>
<version>2.0.17</version>
</dependency>-->
</dependencies>
<properties>
<maven.compiler.source>21</maven.compiler.source>
<maven.compiler.target>21</maven.compiler.target>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
</properties>
</project>
接下来把onnx-bge-small文件夹整个复制到项目目录的resources文件夹下,目录结构如下
src/main/resources/onnx-bge-small/config.json
src/main/resources/onnx-bge-small/model.onnx
src/main/resources/onnx-bge-small/special_tokens_map.json
src/main/resources/onnx-bge-small/tokenizer.json
src/main/resources/onnx-bge-small/tokenizer_config.json
src/main/resources/onnx-bge-small/vocab.txt
接下来我们创建两个类BgeEmbeddingTranslator和BgeSmallEmbedder,代码原理在此不做赘述,需要知道原理的化可以去了解java的djl框架以及深度学习transformer架构。
package com.mocrash.model;
import ai.djl.huggingface.tokenizers.Encoding;
import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import java.io.IOException;
public class BgeEmbeddingTranslator implements Translator<String, float[]> {
private final boolean isQuery;
private static final String QUERY_INSTRUCTION = "Represent this sentence for searching relevant passages: ";
public BgeEmbeddingTranslator(boolean isQuery) {
this.isQuery = isQuery;
}
@Override
public NDList processInput(TranslatorContext ctx, String input) throws IOException {
try (HuggingFaceTokenizer tokenizer = HuggingFaceTokenizer.newInstance(ctx.getModel().getModelPath())) {
String textToEncode = isQuery ? QUERY_INSTRUCTION + input : input; // ✅ 关键修改!
Encoding encoding = tokenizer.encode(textToEncode);
NDManager manager = ctx.getNDManager();
NDArray inputIds = manager.create(encoding.getIds());
NDArray attentionMask = manager.create(encoding.getAttentionMask());
NDArray tokenTypeIds = manager.create(encoding.getTypeIds());
return new NDList(inputIds, attentionMask, tokenTypeIds);
}
}
@Override
public float[] processOutput(TranslatorContext ctx, NDList list) {
NDArray logits = list.singletonOrThrow();
Shape shape = logits.getShape();
NDArray clsVector;
if (shape.dimension() == 2) {
clsVector = logits.get(0);
} else if (shape.dimension() == 3) {
clsVector = logits.get(0).get(0);
} else {
throw new IllegalArgumentException("Unsupported output shape: " + shape);
}
// 确保 L2 归一化
return clsVector.normalize(2, -1).toFloatArray();
}
}
package com.mocrash.model;
import ai.djl.Model;
import ai.djl.inference.Predictor;
import ai.djl.repository.zoo.Criteria;
import ai.djl.translate.TranslateException;
import java.io.IOException;
import java.io.InputStream;
import java.net.URL;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.StandardCopyOption;
public class BgeSmallEmbedder {
private final Model model;
private final BgeEmbeddingTranslator queryTranslator;
private final BgeEmbeddingTranslator passageTranslator;
public BgeSmallEmbedder() throws Exception {
Path modelDir = extractResourceModel("onnx-bge-small");
// 加载 ONNX 模型(指定 engine)
this.model = Model.newInstance("bge-small", "OnnxRuntime");
this.model.load(modelDir);
// 创建两个 translator
this.queryTranslator = new BgeEmbeddingTranslator(true);
this.passageTranslator = new BgeEmbeddingTranslator(false);
}
public void close() {
model.close();
}
// 计算两个向量的余弦相似度(已归一化)
public float cosineSimilarity(float[] a, float[] b) {
double dot = 0.0, normA = 0.0, normB = 0.0;
for (int i = 0; i < a.length; i++) {
dot += a[i] * b[i];
normA += a[i] * a[i];
normB += b[i] * b[i];
}
return (float) (dot / (Math.sqrt(normA) * Math.sqrt(normB)));
}
public float[] embedQuery(String query) throws TranslateException {
try (Predictor<String, float[]> predictor = model.newPredictor(queryTranslator)) {
return predictor.predict(query);
}
}
public float[] embedPassage(String passage) throws TranslateException {
try (Predictor<String, float[]> predictor = model.newPredictor(passageTranslator)) {
return predictor.predict(passage);
}
}
public static Path extractResourceModel(String modelName) throws IOException {
Path tempDir = Files.createTempDirectory("djl-model-");
Path modelDir = tempDir.resolve(modelName);
Files.createDirectories(modelDir);
String[] files = {
"model.onnx",
"config.json",
"special_tokens_map.json",
"tokenizer.json",
"tokenizer_config.json",
"vocab.txt"
};
ClassLoader cl = Thread.currentThread().getContextClassLoader();
for (String file : files) {
String resourcePath = modelName + "/" + file;
URL url = cl.getResource(resourcePath);
if (url == null) {
throw new IOException("Resource not found: " + resourcePath);
}
try (InputStream in = url.openStream()) {
Files.copy(in, modelDir.resolve(file), StandardCopyOption.REPLACE_EXISTING);
}
}
return modelDir;
}
}
准备用于搜索的文档
麻黄适用于风寒感冒无汗、咳嗽气喘、风水浮肿。
桂枝适用于风寒表虚证、肩背肢节酸痛、胸痹心痛。
紫苏适用于风寒感冒兼气滞胸闷、妊娠呕吐、鱼蟹中毒。
生姜适用于风寒感冒轻症、胃寒呕吐、肺寒咳嗽。
荆芥适用于外感表证、麻疹不透、疮疡初起兼表证。
防风适用于外感风寒或风湿痹痛、破伤风、肠风下血。
羌活适用于风寒湿痹、上半身疼痛、太阳头痛。
白芷适用于阳明头痛、鼻渊鼻塞、牙痛、皮肤瘙痒。
细辛适用于少阴头痛、寒饮咳喘、风冷牙痛。
藁本适用于巅顶头痛、风寒湿痹、腹痛泄泻。
苍耳子适用于鼻渊头痛、风疹瘙痒、湿痹拘挛。
辛夷适用于鼻塞流涕、鼻渊头痛、过敏性鼻炎。
薄荷适用于风热感冒、头痛目赤、咽喉肿痛、肝郁气滞。
牛蒡子适用于风热感冒、咳嗽痰多、痄腮喉痹、痈肿疮毒。
蝉蜕适用于风热感冒、咽痛音哑、麻疹不透、小儿惊风。
桑叶适用于风热感冒、肺热燥咳、目赤昏花、肝阳眩晕。
这样我们算是把工具类搭好了,在Springboot中我们可以把这个工具类注册为bean对象(因为加载模型的耗时比较旧,每次调用都重新加载会很慢)。用法参考以下测试方法。
package com.mocrash;
import com.mocrash.model.BgeSmallEmbedder;
import java.io.FileInputStream;
import java.util.ArrayList;
import java.util.List;
import java.util.Scanner;
//TIP 要<b>运行</b>代码,请按 <shortcut actionId="Run"/> 或
// 点击装订区域中的 <icon src="AllIcons.Actions.Execute"/> 图标。
public class Main {
// 方便测试
public static void main(String[] args) throws Exception {
System.out.println("稍等,模型加载中。。。。");
BgeSmallEmbedder embedder = new BgeSmallEmbedder();
System.out.println("模型加载成功。。。。");
System.out.println("稍等,中药库加载中。。。。");
FileInputStream inputStream = new FileInputStream("src/main/resources/中药库.txt");
Scanner fileScan = new Scanner(inputStream);
List<float[]> passageList = new ArrayList<>();
List<String> textList = new ArrayList<>();
while (fileScan.hasNextLine()){
String passage = fileScan.nextLine();
float[] vector = embedder.embedPassage(passage);
passageList.add(vector);
textList.add(passage);
}
System.out.println("稍等,中药库加载完成。。。。");
Scanner scanner = new Scanner(System.in);
while (true){
System.out.print("请输入查询(Query):");
String query = scanner.nextLine();
float[] queryVec = embedder.embedQuery(query);
float maxSimilarity = 0;
String maxText = "";
for (int i = 0; i < passageList.size(); i++) {
float[] passVec = passageList.get(i);
float curSim = embedder.cosineSimilarity(queryVec, passVec);
if (curSim > maxSimilarity){
maxText = textList.get(i);
maxSimilarity = curSim;
}
}
System.out.printf("查到得分%.4f的中药:\n %s\n", maxSimilarity, maxText);
}
}
}

浙公网安备 33010602011771号