问答大模型技术方案算法实现-熵权法融合算法 + 交叉编码器重排算法

算法核心特点

熵权法权重计算:动态计算向量检索和关键词检索的权重,避免固定权重的主观性

交叉编码器重排:使用预训练的交叉编码器模型进行精排,提高相关性判断准确性

两阶段处理:先融合(召回优化),再重排(精度提升)

灵活的权重调整:支持固定权重、熵权法动态权重等多种策略

完整的得分归一化:统一不同检索器的得分尺度

模块化设计:各组件可独立使用和替换

python 复制代码
import numpy as np
from typing import List, Dict, Any, Optional, Tuple
from sklearn.metrics.pairwise import cosine_similarity
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import torch
import torch.nn.functional as F
import logging
from dataclasses import dataclass
from collections import defaultdict
import math

# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


@dataclass
class RetrievalResult:
    """检索结果数据类"""
    doc_id: str
    content: str
    vector_score: float = 0.0
    keyword_score: float = 0.0
    semantic_score: float = 0.0
    final_score: float = 0.0
    normalized_vector_score: float = 0.0
    normalized_keyword_score: float = 0.0
    metadata: Dict[str, Any] = None
    
    def __post_init__(self):
        if self.metadata is None:
            self.metadata = {}


class EntropyWeightCalculator:
    """
    熵权法权重计算器
    用于动态计算向量检索和关键词检索的权重
    """
    
    def __init__(self, epsilon: float = 1e-10):
        """
        初始化熵权法计算器
        
        Args:
            epsilon: 防止除零的小常数
        """
        self.epsilon = epsilon
    
    def calculate_weights(
        self,
        vector_scores: np.ndarray,
        keyword_scores: np.ndarray
    ) -> Tuple[float, float]:
        """
        使用熵权法计算两种检索策略的权重
        
        Args:
            vector_scores: 向量检索得分数组
            keyword_scores: 关键词检索得分数组
            
        Returns:
            (向量权重, 关键词权重)
        """
        # 1. 数据标准化(归一化)
        normalized_vector = self._normalize_scores(vector_scores)
        normalized_keyword = self._normalize_scores(keyword_scores)
        
        # 2. 构建评价矩阵
        eval_matrix = np.vstack([normalized_vector, normalized_keyword]).T
        
        # 3. 计算信息熵
        entropy_vector = self._compute_entropy(normalized_vector)
        entropy_keyword = self._compute_entropy(normalized_keyword)
        
        # 4. 计算权重
        total_entropy = entropy_vector + entropy_keyword
        if total_entropy == 0:
            return 0.5, 0.5
        
        weight_vector = (1 - entropy_vector) / (2 - total_entropy)
        weight_keyword = (1 - entropy_keyword) / (2 - total_entropy)
        
        # 5. 归一化权重
        total_weight = weight_vector + weight_keyword
        if total_weight > 0:
            weight_vector /= total_weight
            weight_keyword /= total_weight
        
        # 限制权重范围
        weight_vector = np.clip(weight_vector, 0.1, 0.9)
        weight_keyword = np.clip(weight_keyword, 0.1, 0.9)
        
        return float(weight_vector), float(weight_keyword)
    
    def _normalize_scores(self, scores: np.ndarray) -> np.ndarray:
        """
        归一化得分(Min-Max标准化)
        
        Args:
            scores: 原始得分数组
            
        Returns:
            归一化后的得分数组
        """
        if len(scores) == 0:
            return np.array([])
        
        min_score = np.min(scores)
        max_score = np.max(scores)
        
        if max_score == min_score:
            return np.ones_like(scores) * 0.5
        
        return (scores - min_score) / (max_score - min_score + self.epsilon)
    
    def _compute_entropy(self, scores: np.ndarray) -> float:
        """
        计算信息熵
        
        Args:
            scores: 归一化后的得分数组
            
        Returns:
            信息熵值
        """
        if len(scores) == 0:
            return 1.0
        
        # 计算概率
        total = np.sum(scores) + self.epsilon
        probabilities = scores / total
        
        # 计算熵
        # 处理概率为0的情况
        probabilities = np.clip(probabilities, self.epsilon, 1.0)
        entropy = -np.sum(probabilities * np.log(probabilities))
        
        # 归一化熵值到[0,1]
        max_entropy = np.log(len(scores) + self.epsilon)
        normalized_entropy = entropy / max_entropy if max_entropy > 0 else 1.0
        
        return float(normalized_entropy)


class CrossEncoderReranker:
    """
    交叉编码器重排序器
    使用预训练的交叉编码器模型进行文档重排
    """
    
    def __init__(
        self,
        model_name: str = "BAAI/bge-reranker-base",
        device: str = "cuda" if torch.cuda.is_available() else "cpu",
        max_length: int = 512,
        batch_size: int = 32
    ):
        """
        初始化交叉编码器重排序器
        
        Args:
            model_name: 模型名称
            device: 运行设备
            max_length: 最大输入长度
            batch_size: 批处理大小
        """
        self.model_name = model_name
        self.device = device
        self.max_length = max_length
        self.batch_size = batch_size
        
        logger.info(f"加载交叉编码器模型: {model_name}")
        logger.info(f"使用设备: {device}")
        
        try:
            self.tokenizer = AutoTokenizer.from_pretrained(model_name)
            self.model = AutoModelForSequenceClassification.from_pretrained(model_name)
            self.model.to(device)
            self.model.eval()
            self.is_loaded = True
            logger.info("交叉编码器模型加载成功")
        except Exception as e:
            logger.warning(f"加载交叉编码器模型失败: {e}")
            logger.info("使用简化版重排序器")
            self.is_loaded = False
    
    def rerank(
        self,
        query: str,
        documents: List[Dict[str, Any]],
        top_k: Optional[int] = None
    ) -> List[Dict[str, Any]]:
        """
        对文档进行重排序
        
        Args:
            query: 查询文本
            documents: 文档列表,每个文档包含 'doc_id', 'content', 'score' 等
            top_k: 重排序的文档数量
            
        Returns:
            重排序后的文档列表
        """
        if not documents:
            return []
        
        # 如果模型未加载,使用简化版本
        if not self.is_loaded:
            return self._simple_rerank(query, documents, top_k)
        
        # 限制重排序的文档数量
        rerank_docs = documents[:top_k] if top_k else documents
        
        # 构建查询-文档对
        pairs = []
        for doc in rerank_docs:
            pairs.append((query, doc.get('content', '')))
        
        # 批量计算相关性得分
        scores = self._compute_relevance_scores(pairs)
        
        # 更新文档得分
        for doc, score in zip(rerank_docs, scores):
            doc['semantic_score'] = float(score)
            doc['rerank_score'] = float(score)
        
        # 按重排序得分排序
        rerank_docs.sort(key=lambda x: x.get('rerank_score', 0), reverse=True)
        
        return rerank_docs
    
    def _compute_relevance_scores(self, pairs: List[Tuple[str, str]]) -> List[float]:
        """
        计算查询-文档对的相关性得分
        
        Args:
            pairs: (查询, 文档) 对列表
            
        Returns:
            相关性得分列表
        """
        if not pairs:
            return []
        
        all_scores = []
        
        for i in range(0, len(pairs), self.batch_size):
            batch_pairs = pairs[i:i + self.batch_size]
            
            # 编码
            inputs = self.tokenizer(
                batch_pairs,
                padding=True,
                truncation=True,
                max_length=self.max_length,
                return_tensors="pt"
            )
            
            inputs = {k: v.to(self.device) for k, v in inputs.items()}
            
            # 推理
            with torch.no_grad():
                outputs = self.model(**inputs)
                scores = outputs.logits.squeeze(-1).cpu().numpy()
            
            # 应用sigmoid得到概率
            scores = 1 / (1 + np.exp(-scores))
            all_scores.extend(scores.tolist())
        
        return all_scores
    
    def _simple_rerank(
        self,
        query: str,
        documents: List[Dict[str, Any]],
        top_k: Optional[int] = None
    ) -> List[Dict[str, Any]]:
        """
        简化版重排序(基于文本重叠度)
        
        Args:
            query: 查询文本
            documents: 文档列表
            top_k: 重排序的文档数量
            
        Returns:
            重排序后的文档列表
        """
        # 提取查询关键词
        query_tokens = set(self._tokenize_simple(query))
        query_tokens = {t for t in query_tokens if len(t) > 1}
        
        rerank_docs = documents[:top_k] if top_k else documents
        
        for doc in rerank_docs:
            doc_tokens = set(self._tokenize_simple(doc.get('content', '')))
            
            # 计算重叠度
            overlap = len(query_tokens & doc_tokens)
            total = len(query_tokens | doc_tokens) + 1
            
            # 计算语义得分(基于重叠度)
            semantic_score = overlap / total
            
            # 结合原有得分
            original_score = doc.get('score', 0)
            doc['semantic_score'] = semantic_score
            doc['rerank_score'] = original_score * 0.7 + semantic_score * 0.3
        
        # 按重排序得分排序
        rerank_docs.sort(key=lambda x: x.get('rerank_score', 0), reverse=True)
        
        return rerank_docs
    
    def _tokenize_simple(self, text: str) -> List[str]:
        """简单的分词(用于简化版重排序)"""
        # 简单的中文分词(按字符分割)
        return list(text)


class DocumentFusionReranker:
    """
    文档融合与重排器(DFR - Document Fusion & Rerank)
    实现文档融合(熵权法)+ 文档重排(交叉编码器)
    """
    
    def __init__(
        self,
        cross_encoder_model: Optional[CrossEncoderReranker] = None,
        entropy_calculator: Optional[EntropyWeightCalculator] = None,
        weight_vector: float = 0.5,
        weight_keyword: float = 0.5,
        fusion_weight: float = 0.5,
        rerank_weight: float = 0.5
    ):
        """
        初始化文档融合与重排器
        
        Args:
            cross_encoder_model: 交叉编码器重排序器
            entropy_calculator: 熵权法计算器
            weight_vector: 向量检索初始权重
            weight_keyword: 关键词检索初始权重
            fusion_weight: 融合阶段权重
            rerank_weight: 重排阶段权重
        """
        self.cross_encoder = cross_encoder_model or CrossEncoderReranker()
        self.entropy_calculator = entropy_calculator or EntropyWeightCalculator()
        
        self.weight_vector = weight_vector
        self.weight_keyword = weight_keyword
        self.fusion_weight = fusion_weight
        self.rerank_weight = rerank_weight
        
        self.weights_history = []  # 记录权重变化历史
    
    def fuse_and_rerank(
        self,
        query: str,
        vector_results: List[Dict[str, Any]],
        keyword_results: List[Dict[str, Any]],
        top_k: int = 10,
        use_entropy_weight: bool = True,
        use_rerank: bool = True
    ) -> List[Dict[str, Any]]:
        """
        执行完整的文档融合与重排流程
        
        Args:
            query: 查询文本
            vector_results: 向量检索结果
            keyword_results: 关键词检索结果
            top_k: 返回结果数量
            use_entropy_weight: 是否使用熵权法计算权重
            use_rerank: 是否使用重排序
            
        Returns:
            融合与重排后的结果列表
        """
        logger.info(f"开始文档融合与重排,向量结果: {len(vector_results)}个,关键词结果: {len(keyword_results)}个")
        
        # 第一阶段:文档融合
        fused_results = self._document_fusion(
            vector_results,
            keyword_results,
            use_entropy_weight
        )
        
        # 第二阶段:文档重排
        if use_rerank and fused_results:
            reranked_results = self._document_rerank(
                query,
                fused_results,
                top_k
            )
        else:
            reranked_results = fused_results[:top_k]
        
        logger.info(f"融合与重排完成,返回 {len(reranked_results)} 个结果")
        
        return reranked_results
    
    def _document_fusion(
        self,
        vector_results: List[Dict[str, Any]],
        keyword_results: List[Dict[str, Any]],
        use_entropy_weight: bool
    ) -> List[Dict[str, Any]]:
        """
        文档融合阶段
        
        Args:
            vector_results: 向量检索结果
            keyword_results: 关键词检索结果
            use_entropy_weight: 是否使用熵权法
            
        Returns:
            融合后的结果列表
        """
        # 1. 提取得分
        vector_scores = np.array([r.get('score', 0) for r in vector_results])
        keyword_scores = np.array([r.get('score', 0) for r in keyword_results])
        
        # 2. 计算权重
        if use_entropy_weight and len(vector_scores) > 0 and len(keyword_scores) > 0:
            w_vector, w_keyword = self.entropy_calculator.calculate_weights(
                vector_scores,
                keyword_scores
            )
            self.weights_history.append({
                'vector_weight': w_vector,
                'keyword_weight': w_keyword
            })
            logger.info(f"熵权法计算权重: vector={w_vector:.3f}, keyword={w_keyword:.3f}")
        else:
            w_vector = self.weight_vector
            w_keyword = self.weight_keyword
        
        # 3. 归一化得分
        normalized_vector = self._normalize_scores(vector_scores)
        normalized_keyword = self._normalize_scores(keyword_scores)
        
        # 4. 构建文档映射
        doc_map = defaultdict(lambda: {
            'doc_id': '',
            'content': '',
            'vector_score': 0,
            'keyword_score': 0,
            'normalized_vector_score': 0,
            'normalized_keyword_score': 0,
            'metadata': {}
        })
        
        # 添加向量检索结果
        for i, result in enumerate(vector_results):
            doc_id = result.get('doc_id', f'vec_{i}')
            doc_map[doc_id]['doc_id'] = doc_id
            doc_map[doc_id]['content'] = result.get('content', '')
            doc_map[doc_id]['vector_score'] = result.get('score', 0)
            doc_map[doc_id]['normalized_vector_score'] = normalized_vector[i] if i < len(normalized_vector) else 0
            doc_map[doc_id]['metadata'] = result.get('metadata', {})
        
        # 添加关键词检索结果
        for i, result in enumerate(keyword_results):
            doc_id = result.get('doc_id', f'key_{i}')
            if doc_id in doc_map:
                doc_map[doc_id]['keyword_score'] = result.get('score', 0)
                doc_map[doc_id]['normalized_keyword_score'] = normalized_keyword[i] if i < len(normalized_keyword) else 0
            else:
                doc_map[doc_id] = {
                    'doc_id': doc_id,
                    'content': result.get('content', ''),
                    'vector_score': 0,
                    'keyword_score': result.get('score', 0),
                    'normalized_vector_score': 0,
                    'normalized_keyword_score': normalized_keyword[i] if i < len(normalized_keyword) else 0,
                    'metadata': result.get('metadata', {})
                }
        
        # 5. 计算融合得分
        fused_results = []
        for doc_id, info in doc_map.items():
            # 融合得分公式: score_a = a * score_sem + (1-a) * score_lex
            fusion_score = (
                w_vector * info['normalized_vector_score'] +
                w_keyword * info['normalized_keyword_score']
            )
            
            result = RetrievalResult(
                doc_id=doc_id,
                content=info['content'],
                vector_score=info['vector_score'],
                keyword_score=info['keyword_score'],
                final_score=fusion_score,
                normalized_vector_score=info['normalized_vector_score'],
                normalized_keyword_score=info['normalized_keyword_score'],
                metadata=info['metadata']
            )
            
            fused_results.append({
                'doc_id': result.doc_id,
                'content': result.content,
                'score': result.final_score,
                'fusion_score': result.final_score,
                'vector_score': result.vector_score,
                'keyword_score': result.keyword_score,
                'normalized_vector_score': result.normalized_vector_score,
                'normalized_keyword_score': result.normalized_keyword_score,
                'metadata': result.metadata,
                'weights': {
                    'vector': w_vector,
                    'keyword': w_keyword
                }
            })
        
        # 按融合得分排序
        fused_results.sort(key=lambda x: x['score'], reverse=True)
        
        return fused_results
    
    def _document_rerank(
        self,
        query: str,
        fused_results: List[Dict[str, Any]],
        top_k: int
    ) -> List[Dict[str, Any]]:
        """
        文档重排阶段
        
        Args:
            query: 查询文本
            fused_results: 融合后的结果列表
            top_k: 返回结果数量
            
        Returns:
            重排后的结果列表
        """
        # 1. 使用交叉编码器计算语义得分
        reranked_results = self.cross_encoder.rerank(
            query,
            fused_results,
            top_k=top_k * 2  # 重排更多的文档
        )
        
        # 2. 融合融合得分和重排得分
        for result in reranked_results:
            fusion_score = result.get('fusion_score', result.get('score', 0))
            semantic_score = result.get('semantic_score', 0)
            
            # 综合得分: score_final = γ * score_β + (1-γ) * score_a
            final_score = (
                self.rerank_weight * semantic_score +
                (1 - self.rerank_weight) * fusion_score
            )
            
            result['final_score'] = final_score
            result['score'] = final_score  # 更新最终得分
        
        # 3. 按最终得分排序
        reranked_results.sort(key=lambda x: x.get('final_score', 0), reverse=True)
        
        # 4. 返回top-k
        return reranked_results[:top_k]
    
    def _normalize_scores(self, scores: np.ndarray) -> np.ndarray:
        """
        归一化得分
        
        Args:
            scores: 原始得分数组
            
        Returns:
            归一化后的得分数组
        """
        if len(scores) == 0:
            return np.array([])
        
        min_score = np.min(scores)
        max_score = np.max(scores)
        
        if max_score == min_score:
            return np.ones_like(scores) * 0.5
        
        return (scores - min_score) / (max_score - min_score + 1e-10)
    
    def get_weights_statistics(self) -> Dict[str, Any]:
        """
        获取权重统计信息
        
        Returns:
            权重统计字典
        """
        if not self.weights_history:
            return {'total_records': 0}
        
        vector_weights = [w['vector_weight'] for w in self.weights_history]
        keyword_weights = [w['keyword_weight'] for w in self.weights_history]
        
        return {
            'total_records': len(self.weights_history),
            'avg_vector_weight': np.mean(vector_weights),
            'avg_keyword_weight': np.mean(keyword_weights),
            'std_vector_weight': np.std(vector_weights),
            'std_keyword_weight': np.std(keyword_weights),
            'min_vector_weight': np.min(vector_weights),
            'max_vector_weight': np.max(vector_weights)
        }


# 测试代码
def test_fusion_reranker():
    """测试文档融合与重排器"""
    
    # 模拟检索结果
    def generate_mock_results(n: int, base_score: float, variance: float):
        results = []
        for i in range(n):
            results.append({
                'doc_id': f'doc_{i:03d}',
                'content': f'文档内容 {i}:这是一段装备维修相关的文本内容。',
                'score': base_score + np.random.randn() * variance,
                'metadata': {'index': i}
            })
        return results
    
    # 生成模拟数据
    np.random.seed(42)
    vector_results = generate_mock_results(30, 0.7, 0.15)
    keyword_results = generate_mock_results(25, 0.6, 0.20)
    
    # 初始化融合重排器
    reranker = DocumentFusionReranker(
        weight_vector=0.6,
        weight_keyword=0.4,
        fusion_weight=0.5,
        rerank_weight=0.5
    )
    
    # 测试查询
    query = "装备故障排查方法"
    
    # 1. 不使用熵权法,不使用重排
    print("="*60)
    print("测试1: 固定权重 + 无重排")
    print("="*60)
    results1 = reranker.fuse_and_rerank(
        query,
        vector_results,
        keyword_results,
        top_k=10,
        use_entropy_weight=False,
        use_rerank=False
    )
    
    for i, result in enumerate(results1[:5], 1):
        print(f"{i}. Doc: {result['doc_id']}, Score: {result['score']:.4f}")
        print(f"   Vector: {result['vector_score']:.4f}, Keyword: {result['keyword_score']:.4f}")
    
    # 2. 使用熵权法,不使用重排
    print("\n" + "="*60)
    print("测试2: 熵权法 + 无重排")
    print("="*60)
    results2 = reranker.fuse_and_rerank(
        query,
        vector_results,
        keyword_results,
        top_k=10,
        use_entropy_weight=True,
        use_rerank=False
    )
    
    for i, result in enumerate(results2[:5], 1):
        print(f"{i}. Doc: {result['doc_id']}, Score: {result['score']:.4f}")
        print(f"   Vector: {result['vector_score']:.4f}, Keyword: {result['keyword_score']:.4f}")
    
    # 3. 使用熵权法 + 重排
    print("\n" + "="*60)
    print("测试3: 熵权法 + 交叉编码器重排")
    print("="*60)
    results3 = reranker.fuse_and_rerank(
        query,
        vector_results,
        keyword_results,
        top_k=10,
        use_entropy_weight=True,
        use_rerank=True
    )
    
    for i, result in enumerate(results3[:5], 1):
        print(f"{i}. Doc: {result['doc_id']}, Score: {result['score']:.4f}")
        print(f"   Fusion: {result.get('fusion_score', 0):.4f}, Semantic: {result.get('semantic_score', 0):.4f}")
    
    # 打印权重统计
    print("\n" + "="*60)
    print("权重统计信息")
    print("="*60)
    stats = reranker.get_weights_statistics()
    for key, value in stats.items():
        print(f"{key}: {value}")
    
    return reranker, results1, results2, results3


def test_entropy_calculator():
    """测试熵权法计算器"""
    print("\n" + "="*60)
    print("熵权法计算器测试")
    print("="*60)
    
    calculator = EntropyWeightCalculator()
    
    # 测试数据
    vector_scores = np.array([0.9, 0.8, 0.7, 0.6, 0.5])
    keyword_scores = np.array([0.7, 0.6, 0.5, 0.4, 0.3])
    
    print(f"向量得分: {vector_scores}")
    print(f"关键词得分: {keyword_scores}")
    
    w_vector, w_keyword = calculator.calculate_weights(vector_scores, keyword_scores)
    print(f"向量权重: {w_vector:.4f}")
    print(f"关键词权重: {w_keyword:.4f}")
    
    # 测试极端情况
    print("\n测试极端情况:")
    vector_scores2 = np.array([1.0, 1.0, 1.0])
    keyword_scores2 = np.array([0.1, 0.2, 0.3])
    
    w_vector2, w_keyword2 = calculator.calculate_weights(vector_scores2, keyword_scores2)
    print(f"向量得分 (全相同): {vector_scores2}")
    print(f"关键词得分: {keyword_scores2}")
    print(f"向量权重: {w_vector2:.4f}")
    print(f"关键词权重: {w_keyword2:.4f}")


if __name__ == "__main__":
    print("="*60)
    print("文档融合与重排算法测试")
    print("="*60)
    
    # 测试熵权法
    test_entropy_calculator()
    
    # 测试融合重排器
    test_fusion_reranker()
相关推荐
阿里云大数据AI技术2 小时前
DataWorks Data Agent 实战课堂(五):AI 多模态智能数据处理
人工智能·agent
又折桃枝换酒钱2 小时前
SASAV:自主科学分析与可视化智能体(翻译与解读)
人工智能
用户5833339816682 小时前
无代码 RPA 对接 LLM:实现文档 OCR+NLP 智能解析的落地实践
人工智能
SelectDB技术团队2 小时前
天翼云 Iceberg 湖仓一体:Apache Doris / SelectDB 的技术能力与实践
人工智能·知识图谱·apache doris·selectdb
面包龙2 小时前
什么是 AI?从人工智能、机器学习到大语言模型和 Agent
人工智能·程序员·全栈
lbb 小魔仙2 小时前
谁替 AI Agent 记住时间?——从航海钟到智能数据库,三百年时延坍缩史
数据库·人工智能·db
满怀冰雪2 小时前
22-使用 PaddleClas 快速训练图像分类模型
人工智能·python·机器学习·分类·数据挖掘·paddlepaddle
我是慎独2 小时前
人工智能:现代方法读书笔记(三)
人工智能·机器学习