Java SpringBoot集成Hugging Face BAAI bg-reranker-large模型方案咨询
集成BAAI bg-reranker-large到Java SpringBoot应用的可行方案
一、使用Deeplearning4j直接加载模型(可行)
Deeplearning4j是Java生态的深度学习框架,支持加载Hugging Face的Transformers系列模型,无需依赖Python转换。它通过ND4J(Java的数值计算库)处理张量运算,结合Transformers模块实现模型的加载与推理。
具体实现步骤
添加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>加载模型与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); } }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
相关产品推荐
相关产品推荐

