在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

准备

  1. 虽然最终使用模型不依赖python技术栈,但是模型需要从python脚本中导出,所以需要准备一个python3的环境。

  2. 为了方便导出,需要一个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);
        }
    }
}
posted @ 2026-02-03 11:34  全栈BUG师  阅读(137)  评论(1)    收藏  举报