You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Java SpringBoot集成Hugging Face BAAI bg-reranker-large模型方案咨询

集成BAAI bg-reranker-large到Java SpringBoot应用的可行方案

一、使用Deeplearning4j直接加载模型(可行)

Deeplearning4j是Java生态的深度学习框架,支持加载Hugging Face的Transformers系列模型,无需依赖Python转换。它通过ND4J(Java的数值计算库)处理张量运算,结合Transformers模块实现模型的加载与推理。

具体实现步骤

  1. 添加Maven依赖
    在pom.xml中引入Deeplearning4j核心、Transformers模块及ND4J后端:

    <dependencies>
        <dependency>
            <groupId>org.deeplearning4j</groupId>
            <artifactId>deeplearning4j-core</artifactId>
            <version>1.0.0-M2.1</version>
        </dependency>
        <dependency>
            <groupId>org.deeplearning4j</groupId>
            <artifactId>deeplearning4j-transformers</artifactId>
            <version>1.0.0-M2.1</version>
        </dependency>
        <dependency>
            <groupId>org.nd4j</groupId>
            <artifactId>nd4j-native-platform</artifactId>
            <version>1.0.0-M2.1</version>
        </dependency>
    </dependencies>
    
  2. 加载模型与Tokenizer
    先从Hugging Face Hub下载bg-reranker-large的模型文件(包括config.json、pytorch_model.bin等)到本地目录,再通过Deeplearning4j加载:

    import org.deeplearning4j.nn.graph.ComputationGraph;
    import org.deeplearning4j.transformers.BertTokenizer;
    import org.deeplearning4j.transformers.model.BertModel;
    import org.nd4j.linalg.api.ndarray.INDArray;
    import org.nd4j.linalg.factory.Nd4j;
    
    public class RerankerService {
        private BertTokenizer tokenizer;
        private ComputationGraph model;
    
        public RerankerService() {
            String modelDir = "path/to/local/bg-reranker-large";
            tokenizer = BertTokenizer.fromPretrained(modelDir);
            model = BertModel.fromPretrained(modelDir, "bg-reranker-large");
        }
    
        public float getRelevanceScore(String query, String passage) {
            // 构造重排模型标准输入格式
            String input = "[CLS] " + query + " [SEP] " + passage + " [SEP]";
            BertTokenizer.TokenizerResult tokenResult = tokenizer.tokenize(input);
    
            // 转换为模型可接受的张量格式
            INDArray inputIds = Nd4j.create(tokenResult.getInputIds());
            INDArray attentionMask = Nd4j.create(tokenResult.getAttentionMask());
            INDArray tokenTypeIds = Nd4j.create(tokenResult.getTokenTypeIds());
    
            // 执行推理,取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token输出作为相关性得分
            INDArray[] outputs = model.output(inputIds, attentionMask, tokenTypeIds);
            return outputs[0].getRow(0).getColumn(0).getFloat(0);
        }
    }
    
  3. SpringBoot集成
    将RerankerService注册为Spring Bean,在业务逻辑中注入使用:

    import org.springframework.stereotype.Service;
    import java.util.AbstractMap;
    import java.util.List;
    import java.util.Map;
    
    @Service
    public class RerankBusinessService {
        private final RerankerService rerankerService;
    
        public RerankBusinessService(RerankerService rerankerService) {
            this.rerankerService = rerankerService;
        }
    
        public List<String> rerankPassages(String query, List<String> passages) {
            return passages.stream()
                    .map(p -> new AbstractMap.SimpleEntry<>(p, rerankerService.getRelevanceScore(query, p)))
                    .sorted((a, b) -> Float.compare(b.getValue(), a.getValue()))
                    .map(Map.Entry::getKey)
                    .toList();
        }
    }
    

二、其他可行方案

1. 封装Python模型为REST API,Java调用

如果Java端直接加载模型存在兼容性问题,可将bg-reranker-large封装为Python API,SpringBoot通过HTTP请求调用,快速实现集成:

  • Python API示例(FastAPI):
    from fastapi import FastAPI
    from transformers import AutoTokenizer, AutoModel
    import torch
    
    app = FastAPI()
    tokenizer = AutoTokenizer.from_pretrained("BAAI/bg-reranker-large")
    model = AutoModel.from_pretrained("BAAI/bg-reranker-large")
    
    @app.post("/rerank")
    def rerank(query: str, passages: list[str]):
        scored = []
        for passage in passages:
            inputs = tokenizer(query, passage, return_tensors="pt")
            with torch.no_grad():
                outputs = model(**inputs)
            score = outputs.last_hidden_state[:, 0, :].numpy()[0][0]
            scored.append((passage, float(score)))
        scored.sort(key=lambda x: x[1], reverse=True)
        return {"reranked": [p for p, s in scored], "scores": [s for p, s in scored]}
    
  • Java调用示例(使用RestTemplate):
    import org.springframework.web.client.RestTemplate;
    import java.util.List;
    import java.util.Map;
    
    public class RerankApiClient {
        private final RestTemplate restTemplate = new RestTemplate();
        private final String apiUrl = "http://localhost:8000/rerank";
    
        public List<String> getRerankedPassages(String query, List<String> passages) {
            Map<String, Object> request = Map.of("query", query, "passages", passages);
            Map<String, Object> response = restTemplate.postForObject(apiUrl, request, Map.class);
            return (List<String>) response.get("reranked");
        }
    }
    

2. ONNX Runtime for Java(仅需一次Python导出)

虽然你提到ONNX需要Python,但仅需一次Python执行即可完成模型导出,后续Java端可直接加载ONNX模型运行,无需再依赖Python:

  • 导出ONNX模型(仅执行一次):
    from transformers import AutoTokenizer, AutoModel
    from transformers.onnx import export_onnx
    
    model_name = "BAAI/bg-reranker-large"
    tokenizer = AutoTokenizer.from_pretrained(model_name)
    model = AutoModel.from_pretrained(model_name)
    export_onnx(model, tokenizer, output="bg-reranker-large.onnx", opset=13)
    
  • Java端依赖与代码:
    <dependency>
        <groupId>com.microsoft.onnxruntime</groupId>
        <artifactId>onnxruntime</artifactId>
        <version>1.15.1</version>
    </dependency>
    <dependency>
        <groupId>com.huggingface</groupId>
        <artifactId>tokenizers</artifactId>
        <version>0.13.3</version>
    </dependency>
    
    import ai.onnxruntime.*;
    import com.huggingface.tokenizers.AutoTokenizer;
    import com.huggingface.tokenizers.Encoding;
    import java.util.HashMap;
    import java.util.Map;
    
    public class ONNXReranker {
        private OrtEnvironment env;
        private OrtSession session;
        private AutoTokenizer tokenizer;
    
        public ONNXReranker() throws OrtException {
            env = OrtEnvironment.getEnvironment();
            session = env.createSession("bg-reranker-large.onnx", new OrtSession.SessionOptions());
            tokenizer = AutoTokenizer.fromPretrained("BAAI/bg-reranker-large");
        }
    
        public float calculateScore(String query, String passage) {
            Encoding encoding = tokenizer.encode(query, passage);
            Map<String, OnnxTensor> inputs = new HashMap<>();
            try {
                inputs.put("input_ids", OnnxTensor.createTensor(env, encoding.getIds()));
                inputs.put("attention_mask", OnnxTensor.createTensor(env, encoding.getAttentionMask()));
                inputs.put("token_type_ids", OnnxTensor.createTensor(env, encoding.getTypeIds()));
    
                try (OrtSession.Result result = session.run(inputs)) {
                    float[][] output = (float[][]) result.get(0).getValue();
                    return output[0][0];
                }
            } catch (OrtException e) {
                throw new RuntimeException(e);
            } finally {
                inputs.values().forEach(OnnxTensor::close);
            }
        }
    }
    

内容的提问来源于stack exchange,提问作者chulin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 02:10:58