Java案例如何实现语义搜索?

wen python案例 2

本文目录导读:

Java案例如何实现语义搜索?

  1. Java实现语义搜索的完整指南
  2. 基于嵌入向量的语义搜索(最常用)
  3. 使用预训练模型的实现(推荐)
  4. 使用Elasticsearch的向量搜索
  5. 完整示例:基于Word2Vec的语义搜索
  6. 高级优化方案
  7. 性能优化建议

Java实现语义搜索的完整指南

语义搜索比传统的关键词搜索更智能,它能理解查询的意图和上下文,下面我将介绍几种Java实现语义搜索的方案。

基于嵌入向量的语义搜索(最常用)

核心原理

  • 将文本转换为向量(嵌入)
  • 计算向量之间的余弦相似度
  • 返回相似度最高的结果

实现步骤

步骤1:添加依赖

<dependency>
    <groupId>org.deeplearning4j</groupId>
    <artifactId>deeplearning4j-core</artifactId>
    <version>1.0.0-M2</version>
</dependency>

步骤2:基础实现

import java.util.*;
import java.util.stream.Collectors;
public class SimpleSemanticSearch {
    // 简单的TF-IDF向量化(示例用,实际应使用预训练模型)
    public static double[] vectorize(String text, Map<String, Double> idfMap) {
        String[] words = text.toLowerCase().split("\\s+");
        double[] vector = new double[idfMap.size()];
        Map<String, Long> wordCount = Arrays.stream(words)
            .collect(Collectors.groupingBy(w -> w, Collectors.counting()));
        int index = 0;
        for (String word : idfMap.keySet()) {
            double tf = wordCount.getOrDefault(word, 0L) / (double) words.length;
            double idf = idfMap.get(word);
            vector[index++] = tf * idf;
        }
        return vector;
    }
    // 余弦相似度计算
    public static double cosineSimilarity(double[] vec1, double[] vec2) {
        double dotProduct = 0.0;
        double norm1 = 0.0;
        double norm2 = 0.0;
        for (int i = 0; i < vec1.length; i++) {
            dotProduct += vec1[i] * vec2[i];
            norm1 += vec1[i] * vec1[i];
            norm2 += vec2[i] * vec2[i];
        }
        return dotProduct / (Math.sqrt(norm1) * Math.sqrt(norm2));
    }
}

使用预训练模型的实现(推荐)

使用Sentence-BERT/SBERT

import ai.djl.Application;
import ai.djl.ModelException;
import ai.djl.inference.Predictor;
import ai.djl.modality.nlp.bert.BertTokenizer;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelZoo;
import ai.djl.repository.zoo.ZooModel;
import ai.djl.training.util.ProgressBar;
public class SemanticSearchWithSBERT {
    private ZooModel<String[], float[]> model;
    private Predictor<String[], float[]> predictor;
    public void init() throws ModelException, IOException {
        Criteria<String[], float[]> criteria = Criteria.builder()
                .optApplication(Application.NLP.TEXT_EMBEDDING)
                .setTypes(String[].class, float[].class)
                .optModelUrls("https://resources.djl.ai/models/sentence-transformers/all-MiniLM-L6-v2.zip")
                .optProgress(new ProgressBar())
                .build();
        model = ModelZoo.loadModel(criteria);
        predictor = model.newPredictor();
    }
    public float[] getEmbedding(String text) throws Exception {
        return predictor.predict(new String[]{text});
    }
    public List<SearchResult> search(String query, List<Document> documents) {
        float[] queryVector = getEmbedding(query);
        List<SearchResult> results = new ArrayList<>();
        for (Document doc : documents) {
            float[] docVector = getEmbedding(doc.getText());
            double similarity = cosineSimilarity(queryVector, docVector);
            results.add(new SearchResult(doc, similarity));
        }
        results.sort((a, b) -> Double.compare(b.getScore(), a.getScore()));
        return results;
    }
    private double cosineSimilarity(float[] vec1, float[] vec2) {
        double dot = 0.0;
        double norm1 = 0.0;
        double norm2 = 0.0;
        for (int i = 0; i < vec1.length; i++) {
            dot += vec1[i] * vec2[i];
            norm1 += vec1[i] * vec1[i];
            norm2 += vec2[i] * vec2[i];
        }
        return dot / (Math.sqrt(norm1) * Math.sqrt(norm2));
    }
}
class SearchResult {
    private Document document;
    private double score;
    // 构造函数、getter/setter
}

使用Elasticsearch的向量搜索

import org.elasticsearch.action.search.SearchRequest;
import org.elasticsearch.action.search.SearchResponse;
import org.elasticsearch.client.RestHighLevelClient;
import org.elasticsearch.index.query.QueryBuilders;
import org.elasticsearch.script.Script;
import org.elasticsearch.script.ScriptType;
import org.elasticsearch.search.builder.SearchSourceBuilder;
public class ElasticsearchSemanticSearch {
    private RestHighLevelClient client;
    public List<Document> semanticSearch(String query, String indexName) {
        // 1. 向量化查询(使用外部模型)
        float[] queryVector = convertToVector(query);
        // 2. 构建脚本查询(余弦相似度)
        SearchSourceBuilder sourceBuilder = new SearchSourceBuilder();
        sourceBuilder.query(QueryBuilders.scriptScoreQuery(
            QueryBuilders.matchAllQuery(),
            new Script(ScriptType.INLINE, "painless",
                "cosineSimilarity(params.queryVector, 'vector_field') + 1.0",
                Collections.singletonMap("queryVector", queryVector))
        ));
        // 3. 执行搜索
        SearchRequest searchRequest = new SearchRequest(indexName);
        searchRequest.source(sourceBuilder);
        try {
            SearchResponse response = client.search(searchRequest, RequestOptions.DEFAULT);
            return parseResponse(response);
        } catch (IOException e) {
            e.printStackTrace();
            return Collections.emptyList();
        }
    }
}

完整示例:基于Word2Vec的语义搜索

import org.deeplearning4j.models.embeddings.loader.WordVectorSerializer;
import org.deeplearning4j.models.word2vec.Word2Vec;
import org.nd4j.linalg.api.ndarray.INDArray;
import org.nd4j.linalg.ops.transforms.Transforms;
public class Word2VecSemanticSearch {
    private Word2Vec word2vec;
    public void loadModel(String modelPath) throws Exception {
        word2vec = WordVectorSerializer.readWord2VecModel(modelPath);
    }
    public double[] sentenceToVector(String sentence) {
        String[] words = sentence.toLowerCase().split("\\s+");
        double[] sumVector = new double[100]; // Word2Vec维度
        int count = 0;
        for (String word : words) {
            if (word2vec.hasWord(word)) {
                double[] wordVector = word2vec.getWordVector(word);
                for (int i = 0; i < wordVector.length; i++) {
                    sumVector[i] += wordVector[i];
                }
                count++;
            }
        }
        // 平均向量
        if (count > 0) {
            for (int i = 0; i < sumVector.length; i++) {
                sumVector[i] /= count;
            }
        }
        return sumVector;
    }
    public List<SearchResult> search(String query, List<Document> documents) {
        double[] queryVector = sentenceToVector(query);
        List<SearchResult> results = new ArrayList<>();
        for (Document doc : documents) {
            double[] docVector = sentenceToVector(doc.getText());
            double similarity = cosineSimilarity(queryVector, docVector);
            results.add(new SearchResult(doc, similarity));
        }
        results.sort((a, b) -> Double.compare(b.getScore(), a.getScore()));
        return results;
    }
    public static double cosineSimilarity(double[] vector1, double[] vector2) {
        if (vector1.length != vector2.length) {
            throw new IllegalArgumentException("Vectors must have same length");
        }
        double dotProduct = 0.0;
        double norm1 = 0.0;
        double norm2 = 0.0;
        for (int i = 0; i < vector1.length; i++) {
            dotProduct += vector1[i] * vector2[i];
            norm1 += vector1[i] * vector1[i];
            norm2 += vector2[i] * vector2[i];
        }
        if (norm1 == 0 || norm2 == 0) {
            return 0.0;
        }
        return dotProduct / (Math.sqrt(norm1) * Math.sqrt(norm2));
    }
}

高级优化方案

使用近似最近邻搜索(ANN)

import io.pinecone.PineconeClient;
import io.pinecone.PineconeIndex;
public class ANNBasedSearch {
    private PineconeClient pineconeClient;
    public void init() {
        pineconeClient = new PineconeClient.Builder("your-api-key").build();
    }
    public List<SearchResult> searchWithANN(float[] queryVector, int topK) {
        // 使用Pinecone或其他向量数据库
        QueryResponse response = pineconeClient.query(
            "index-name", 
            queryVector,
            topK,
            true  // 包含元数据
        );
        return response.getMatches().stream()
            .map(match -> new SearchResult(match.getMetadata(), match.getScore()))
            .collect(Collectors.toList());
    }
}

性能优化建议

批量处理

public class BatchProcessor {
    private ExecutorService executor = Executors.newFixedThreadPool(4);
    public List<SearchResult> batchSearch(String query, List<Document> documents) {
        List<Future<SearchResult>> futures = documents.stream()
            .map(doc -> executor.submit(() -> computeSimilarity(query, doc)))
            .collect(Collectors.toList());
        return futures.stream()
            .map(f -> {
                try { return f.get(); }
                catch (Exception e) { return null; }
            })
            .filter(Objects::nonNull)
            .sorted(Comparator.comparingDouble(SearchResult::getScore).reversed())
            .collect(Collectors.toList());
    }
}

选择哪种方案取决于您的具体需求:

  1. 小规模数据:使用本地Word2Vec或BERT模型
  2. 生产环境:使用Elasticsearch向量搜索或Pinecone等向量数据库
  3. 需要高精度:使用预训练的Sentence-BERT模型
  4. 大规模数据:使用近似最近邻搜索(ANN)

建议从简单的实现开始,根据实际需求逐步优化。

抱歉,评论功能暂时关闭!