Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
weight_rerank.py186 linesDownload Raw Back to rerank
1import math2from collections import Counter3from typing import Optional4 5import numpy as np6 7from core.model_manager import ModelManager8from core.model_runtime.entities.model_entities import ModelType9from core.rag.datasource.keyword.jieba.jieba_keyword_table_handler import JiebaKeywordTableHandler10from core.rag.embedding.cached_embedding import CacheEmbedding11from core.rag.models.document import Document12from core.rag.rerank.entity.weight import VectorSetting, Weights13from core.rag.rerank.rerank_base import BaseRerankRunner14 15 16class WeightRerankRunner(BaseRerankRunner):17    def __init__(self, tenant_id: str, weights: Weights) -> None:18        self.tenant_id = tenant_id19        self.weights = weights20 21    def run(22        self,23        query: str,24        documents: list[Document],25        score_threshold: Optional[float] = None,26        top_n: Optional[int] = None,27        user: Optional[str] = None,28    ) -> list[Document]:29        """30        Run rerank model31        :param query: search query32        :param documents: documents for reranking33        :param score_threshold: score threshold34        :param top_n: top n35        :param user: unique user id if needed36 37        :return:38        """39        docs = []40        doc_id = []41        unique_documents = []42        for document in documents:43            if document.metadata["doc_id"] not in doc_id:44                doc_id.append(document.metadata["doc_id"])45                docs.append(document.page_content)46                unique_documents.append(document)47 48        documents = unique_documents49 50        rerank_documents = []51        query_scores = self._calculate_keyword_score(query, documents)52 53        query_vector_scores = self._calculate_cosine(self.tenant_id, query, documents, self.weights.vector_setting)54        for document, query_score, query_vector_score in zip(documents, query_scores, query_vector_scores):55            # format document56            score = (57                self.weights.vector_setting.vector_weight * query_vector_score58                + self.weights.keyword_setting.keyword_weight * query_score59            )60            if score_threshold and score < score_threshold:61                continue62            document.metadata["score"] = score63            rerank_documents.append(document)64        rerank_documents = sorted(rerank_documents, key=lambda x: x.metadata["score"], reverse=True)65        return rerank_documents[:top_n] if top_n else rerank_documents66 67    def _calculate_keyword_score(self, query: str, documents: list[Document]) -> list[float]:68        """69        Calculate BM25 scores70        :param query: search query71        :param documents: documents for reranking72 73        :return:74        """75        keyword_table_handler = JiebaKeywordTableHandler()76        query_keywords = keyword_table_handler.extract_keywords(query, None)77        documents_keywords = []78        for document in documents:79            # get the document keywords80            document_keywords = keyword_table_handler.extract_keywords(document.page_content, None)81            document.metadata["keywords"] = document_keywords82            documents_keywords.append(document_keywords)83 84        # Counter query keywords(TF)85        query_keyword_counts = Counter(query_keywords)86 87        # total documents88        total_documents = len(documents)89 90        # calculate all documents' keywords IDF91        all_keywords = set()92        for document_keywords in documents_keywords:93            all_keywords.update(document_keywords)94 95        keyword_idf = {}96        for keyword in all_keywords:97            # calculate include query keywords' documents98            doc_count_containing_keyword = sum(1 for doc_keywords in documents_keywords if keyword in doc_keywords)99            # IDF100            keyword_idf[keyword] = math.log((1 + total_documents) / (1 + doc_count_containing_keyword)) + 1101 102        query_tfidf = {}103 104        for keyword, count in query_keyword_counts.items():105            tf = count106            idf = keyword_idf.get(keyword, 0)107            query_tfidf[keyword] = tf * idf108 109        # calculate all documents' TF-IDF110        documents_tfidf = []111        for document_keywords in documents_keywords:112            document_keyword_counts = Counter(document_keywords)113            document_tfidf = {}114            for keyword, count in document_keyword_counts.items():115                tf = count116                idf = keyword_idf.get(keyword, 0)117                document_tfidf[keyword] = tf * idf118            documents_tfidf.append(document_tfidf)119 120        def cosine_similarity(vec1, vec2):121            intersection = set(vec1.keys()) & set(vec2.keys())122            numerator = sum(vec1[x] * vec2[x] for x in intersection)123 124            sum1 = sum(vec1[x] ** 2 for x in vec1)125            sum2 = sum(vec2[x] ** 2 for x in vec2)126            denominator = math.sqrt(sum1) * math.sqrt(sum2)127 128            if not denominator:129                return 0.0130            else:131                return float(numerator) / denominator132 133        similarities = []134        for document_tfidf in documents_tfidf:135            similarity = cosine_similarity(query_tfidf, document_tfidf)136            similarities.append(similarity)137 138        # for idx, similarity in enumerate(similarities):139        #     print(f"Document {idx + 1} similarity: {similarity}")140 141        return similarities142 143    def _calculate_cosine(144        self, tenant_id: str, query: str, documents: list[Document], vector_setting: VectorSetting145    ) -> list[float]:146        """147        Calculate Cosine scores148        :param query: search query149        :param documents: documents for reranking150 151        :return:152        """153        query_vector_scores = []154 155        model_manager = ModelManager()156 157        embedding_model = model_manager.get_model_instance(158            tenant_id=tenant_id,159            provider=vector_setting.embedding_provider_name,160            model_type=ModelType.TEXT_EMBEDDING,161            model=vector_setting.embedding_model_name,162        )163        cache_embedding = CacheEmbedding(embedding_model)164        query_vector = cache_embedding.embed_query(query)165        for document in documents:166            # calculate cosine similarity167            if "score" in document.metadata:168                query_vector_scores.append(document.metadata["score"])169            else:170                # transform to NumPy171                vec1 = np.array(query_vector)172                vec2 = np.array(document.vector)173 174                # calculate dot product175                dot_product = np.dot(vec1, vec2)176 177                # calculate norm178                norm_vec1 = np.linalg.norm(vec1)179                norm_vec2 = np.linalg.norm(vec2)180 181                # calculate cosine similarity182                cosine_sim = dot_product / (norm_vec1 * norm_vec2)183                query_vector_scores.append(cosine_sim)184 185        return query_vector_scores186