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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 17:45:17