如何用Deep Java Library将all-MiniLM-L6-v2预训练模型用作Cross Encoder
使用DJL实现Cross Encoder文本对相似度计算(用于文档重排序)
下面是修改后的代码,适配Cross Encoder的使用场景,可直接计算查询语句与检索到的文档之间的相似度分数,实现聊天机器人的文档重排序:
package ai.djl.examples.inference.nlp; import ai.djl.ModelException; import ai.djl.huggingface.translator.CrossEncoderTranslator; import ai.djl.inference.Predictor; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ZooModel; import ai.djl.training.util.ProgressBar; import ai.djl.translate.TranslateException; import java.io.IOException; import java.util.ArrayList; import java.util.Comparator; import java.util.List; public class CrossEncoderReranking { private CrossEncoderReranking() {} public static void main(String[] args) throws IOException, ModelException, TranslateException { // 聊天机器人的用户查询 String query = "如何在DJL中使用Cross Encoder?"; // 检索到的候选文档列表 List<String> retrievedDocuments = List.of( "DJL支持多种HuggingFace模型,包括文本嵌入和Cross Encoder", "文本嵌入模型用于生成单个文本的向量表示", "Cross Encoder可以直接计算两个文本之间的相似度分数,适合重排序", "DJL的PyTorch引擎支持大部分Transformer模型" ); // 构建Cross Encoder的Criteria Criteria<String[], Float> criteria = Criteria.builder() // 输入为文本对数组:[query, document],输出为相似度分数 .setTypes(String[].class, Float.class) // 使用专为检索重排序训练的Cross Encoder模型 .optModelUrls("djl://ai.djl.huggingface.pytorch/sentence-transformers/msmarco-MiniLM-L-6-v2") .optEngine("PyTorch") // 使用CrossEncoderTranslator处理文本对的预处理和后处理 .optTranslator(CrossEncoderTranslator.builder().build()) .optProgress(new ProgressBar()) .build(); try (ZooModel<String[], Float> model = criteria.loadModel(); Predictor<String[], Float> predictor = model.newPredictor()) { // 存储文档与对应分数的列表 List<DocumentScore> documentScores = new ArrayList<>(); // 遍历每个检索到的文档,计算与查询的相似度分数 for (String doc : retrievedDocuments) { String[] inputPair = new String[]{query, doc}; Float score = predictor.predict(inputPair); documentScores.add(new DocumentScore(doc, score)); } // 按相似度分数降序排序,实现重排序 documentScores.sort(Comparator.comparing(DocumentScore::getScore).reversed()); // 输出重排序结果 System.out.println("重排序后的文档(按相似度从高到低):"); for (int i = 0; i < documentScores.size(); i++) { DocumentScore ds = documentScores.get(i); System.out.printf("%d. 分数: %.4f | 内容: %s%n", i+1, ds.getScore(), ds.getDocument()); } } } // 辅助类,用于存储文档和对应的相似度分数 static class DocumentScore { private final String document; private final Float score; public DocumentScore(String document, Float score) { this.document = document; this.score = score; } public String getDocument() { return document; } public Float getScore() { return score; } } }
关键修改说明:
- 模型替换:使用
msmarco-MiniLM-L-6-v2模型,这是专门为检索场景训练的Cross Encoder,比通用模型更适合文档重排序任务。 - 输入输出类型调整:输入改为
String[](存储查询和文档的文本对),输出改为Float(直接输出相似度分数)。 - Translator更换:使用
CrossEncoderTranslator替代文本嵌入的翻译器,它会自动处理文本对的拼接、tokenization以及分数输出。 - 重排序逻辑:新增了文档与分数的存储类,计算所有候选文档的分数后按降序排序,直接得到重排序结果。
内容的提问来源于stack exchange,提问作者Azeem Waqar
相关产品推荐
相关产品推荐

