企业本地知识库助手 大模型+本地知识库

作者:

在

打造企业本地知识库助手,结合了企业私有的知识库,又利用了大模型的生成能力。

整体流程:

1、使用 embeddings模型将企业知识文档向量化存入向量数据库Faiss中

2、接收用户问题,从本地知识库中查找相关性排名前N的知识块

3、让大模型依据本地知识回答用户问题

embeddings模型

embeddings文本向量化模型选择,为了减少下载与启动时长,快速跑起demo,我使用的是iic/nlp_gte_sentence-embedding_chinese-base模型,大概几百MB,下载比较快。可以对比通义千问系列其他模型,文件比较大,可能语文理解会更好。当量,也可以使用直接调用embeddings在线api,避免本地部署。

大模型大模型对机器的配置要求更高,可直接调用大模型服务,或者运行小尺寸的大模型。上代码代码我在本地调试运行过,可直接copy过去尝试一下。可稍微完善一下,就是一个企业本地知识库的智能助手。

import os
import sys
import time
import logging
import argparse
from typing import List, Tuple, Dict, Optional
from langchain_community.llms import Ollama
from langchain_openai import ChatOpenAI

# 配置日志
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)

# 检查必要的依赖库
try:
    from PyPDF2 import PdfReader
    from langchain.chains.question_answering import load_qa_chain
    from langchain_openai import OpenAI
    from langchain_community.callbacks.manager import get_openai_callback
    from langchain.text_splitter import RecursiveCharacterTextSplitter
    from langchain_community.embeddings import HuggingFaceEmbeddings
    from langchain_community.vectorstores import FAISS
    from modelscope import snapshot_download, AutoModel, AutoTokenizer
    import torch
    import torch.nn.functional as F
    from langchain.embeddings.base import Embeddings
except ImportError as e:
    print(f"错误: 缺少必要的依赖库: {str(e)}")
    print("请安装所需依赖: pip install PyPDF2 langchain langchain-openai langchain-community faiss-cpu modelscope torch")
    sys.exit(1)

# 自定义 ModelScope 嵌入类
class ModelScopeEmbeddings(Embeddings):
    """使用 ModelScope 加载的模型进行文本嵌入"""

    # 类变量用于缓存已加载的模型
    _models_cache = {}

    def __init__(self, model_id: str = "iic/nlp_gte_sentence-embedding_chinese-base", 
                 device: str = None, batch_size: int = 32, use_query_prefix: bool = True):
        """
        初始化 ModelScope 嵌入模型

        参数:
            model_id: ModelScope 模型ID
            device: 计算设备,如 'cuda:0', 'cpu' 等,默认为自动选择
            batch_size: 批处理大小,用于处理大量文档
            use_query_prefix: 是否为查询添加前缀
        """
        self.model_id = model_id
        self.batch_size = batch_size
        self.use_query_prefix = use_query_prefix

        # 设置设备
        if device is None:
            self.device = "cuda" if torch.cuda.is_available() else "cpu"
        else:
            self.device = device

        # 缓存键
        cache_key = f"{model_id}_{self.device}"

        try:
            # 检查模型是否已经加载
            if cache_key in self._models_cache:
                logger.info(f"使用缓存中的模型: {model_id}")
                self.tokenizer, self.model = self._models_cache[cache_key]
            else:
                # 检查模型是否已经下载
                try:
                    # 尝试直接加载本地模型
                    self.model_dir = snapshot_download(model_id, cache_dir=os.path.join(os.path.expanduser("~"), ".cache", "modelscope"))
                    logger.info(f"模型已存在于本地缓存: {self.model_dir}")
                except Exception as e:
                    logger.info(f"本地未找到模型,正在从 ModelScope 下载模型 {model_id}...")
                    try:
                        self.model_dir = snapshot_download(model_id)
                        logger.info(f"模型已下载到: {self.model_dir}")
                    except Exception as download_error:
                        logger.error(f"下载模型时出错: {str(download_error)}")
                        raise ValueError(f"无法下载模型 {model_id}: {str(download_error)}")

                # 加载模型和分词器
                logger.info(f"正在加载模型到 {self.device} 设备...")
                try:
                    self.tokenizer = AutoTokenizer.from_pretrained(self.model_dir)
                    self.model = AutoModel.from_pretrained(self.model_dir)
                    self.model.to(self.device)

                    # 缓存模型
                    self._models_cache[cache_key] = (self.tokenizer, self.model)
                    logger.info(f"模型已加载并缓存")
                except Exception as load_error:
                    logger.error(f"加载模型时出错: {str(load_error)}")
                    raise ValueError(f"无法加载模型: {str(load_error)}")
        except Exception as e:
            logger.error(f"初始化 ModelScope 嵌入模型时出错: {str(e)}")
            raise ValueError(f"初始化失败: {str(e)}")

    def _mean_pooling(self, model_output, attention_mask):
        """平均池化获取句子嵌入"""
        token_embeddings = model_output[0]
        input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
        return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(input_mask_expanded.sum(1), min=1e-9)

    def embed_documents(self, texts: List[str]) -> List[List[float]]:
        """
        为文档生成嵌入向量

        参数:
            texts: 文档文本列表

        返回:
            embeddings: 嵌入向量列表
        """
        if not texts:
            return []

        try:
            # 分批处理文档
            all_embeddings = []
            for i in range(0, len(texts), self.batch_size):
                batch_texts = texts[i:i+self.batch_size]

                # 对文本进行编码
                encoded_input = self.tokenizer(
                    batch_texts, 
                    padding=True, 
                    truncation=True, 
                    max_length=512,  # 添加最大长度限制
                    return_tensors='pt'
                ).to(self.device)

                # 生成嵌入
                with torch.no_grad():
                    model_output = self.model(**encoded_input)

                # 池化并归一化
                sentence_embeddings = self._mean_pooling(model_output, encoded_input['attention_mask'])
                normalized_embeddings = F.normalize(sentence_embeddings, p=2, dim=1)

                # 添加到结果列表
                all_embeddings.extend(normalized_embeddings.cpu().tolist())

            return all_embeddings
        except Exception as e:
            logger.error(f"生成文档嵌入时出错: {str(e)}")
            # 在出错时返回零向量,避免程序崩溃
            if len(texts) > 0:
                try:
                    # 尝试获取嵌入维度
                    dummy_output = self.embed_documents([texts[0]])
                    dim = len(dummy_output[0])
                    return [[0.0] * dim] * len(texts)
                except:
                    # 如果无法确定维度,使用默认维度
                    return [[0.0] * 768] * len(texts)
            return []

    def embed_query(self, text: str) -> List[float]:
        """
        为查询生成嵌入向量

        参数:
            text: 查询文本

        返回:
            embedding: 嵌入向量
        """
        if not text:
            # 返回零向量
            return [0.0] * 768

        try:
            # 为查询添加提示词(如果启用)
            query_text = f"查询:{text}" if self.use_query_prefix else text
            return self.embed_documents([query_text])[0]
        except Exception as e:
            logger.error(f"生成查询嵌入时出错: {str(e)}")
            # 返回零向量
            return [0.0] * 768

class PDFChatBot:
    """PDF聊天机器人类,用于处理PDF文档并回答相关问题"""

    def __init__(self, pdf_path: str, api_key: str, base_url: str = "https://dashscope.aliyuncs.com/compatible-mode/v1", 
                 model_id: str = "iic/nlp_gte_sentence-embedding_chinese-base", 
                 device: str = None, use_query_prefix: bool = True,
                 llm_model_name: str = "deepseek-r1:1.5b"):
        """
        初始化PDF聊天机器人

        参数:
            pdf_path: PDF文件路径
            api_key: OpenAI API密钥
            base_url: API基础URL
            model_id: ModelScope模型ID,用于文本嵌入
            device: 计算设备,如 'cuda:0', 'cpu' 等,默认为自动选择
            use_query_prefix: 是否为查询添加前缀
        """
        self.pdf_path = pdf_path
        self.api_key = api_key
        self.base_url = base_url
        self.model_id = model_id
        self.device = device
        self.use_query_prefix = use_query_prefix
        self.llm_model_name = llm_model_name
        self.knowledge_base = None
        self.page_count = 0

        # 检查PDF文件是否存在
        if not os.path.exists(pdf_path):
            raise FileNotFoundError(f"PDF文件不存在: {pdf_path}")

        logger.info(f"初始化PDF聊天机器人,使用文件: {pdf_path}")

    def load_pdf(self) -> None:
        """加载PDF文件并创建知识库"""
        start_time = time.time()
        logger.info(f"开始加载PDF文件: {self.pdf_path}")

        try:
            # 读取PDF文件
            pdf_reader = PdfReader(self.pdf_path)
            self.page_count = len(pdf_reader.pages)
            logger.info(f"PDF文件共有 {self.page_count} 页")

            # 提取文本和页码信息
            text, page_numbers = self._extract_text_with_page_numbers(pdf_reader)
            logger.info(f"提取的文本长度: {len(text)} 个字符")

            if not text:
                raise ValueError("无法从PDF中提取文本,请检查PDF文件是否有效")

            # 处理文本并创建知识库
            self.knowledge_base = self._process_text_with_splitter(text, page_numbers)

            elapsed_time = time.time() - start_time
            logger.info(f"PDF加载和知识库创建完成,耗时: {elapsed_time:.2f} 秒")
        except Exception as e:
            logger.error(f"加载PDF文件时出错: {str(e)}")
            raise

    def _extract_text_with_page_numbers(self, pdf) -> Tuple[str, List[int]]:
        """
        从PDF中提取文本并记录每页文本对应的页码

        参数:
            pdf: PDF文件对象

        返回:
            text: 提取的文本内容
            page_numbers: 每页文本对应的页码列表
        """
        text = ""
        page_texts = []  # 存储每页的文本
        page_numbers = []  # 存储每页的页码

        for page_number, page in enumerate(pdf.pages, start=1):
            try:
                extracted_text = page.extract_text()
                if extracted_text:
                    text += extracted_text + "\n\n"  # 添加页面分隔符
                    page_texts.append(extracted_text)
                    page_numbers.append(page_number)
                else:
                    logger.warning(f"第 {page_number} 页未找到文本")
            except Exception as e:
                logger.error(f"提取第 {page_number} 页文本时出错: {str(e)}")

        return text, page_numbers

    def _process_text_with_splitter(self, text: str, page_numbers: List[int]) -> FAISS:
        """
        处理文本并创建向量存储

        参数:
            text: 提取的文本内容
            page_numbers: 每页文本对应的页码列表

        返回:
            knowledgeBase: 基于FAISS的向量存储对象
        """
        try:
            # 创建文本分割器,用于将长文本分割成小块
            text_splitter = RecursiveCharacterTextSplitter(
                separators=["\n\n", "\n", ".", " ", ""],
                chunk_size=1000,
                chunk_overlap=200,
                length_function=len,
            )

            # 分割文本
            chunks = text_splitter.split_text(text)
            logger.info(f"文本被分割成 {len(chunks)} 个块")

            try:
                # 创建嵌入模型 - 使用ModelScope的中文嵌入模型
                logger.info(f"正在初始化 ModelScope 嵌入模型: {self.model_id}...")
                embeddings = ModelScopeEmbeddings(
                    model_id=self.model_id,
                    device=self.device,
                    use_query_prefix=self.use_query_prefix
                )

                # 从文本块创建知识库
                knowledge_base = FAISS.from_texts(chunks, embeddings)
                logger.info("已从文本块创建知识库")

                # 为每个文本块分配页码
                knowledge_base.page_info = self._assign_page_numbers_to_chunks(text, chunks, page_numbers)

                return knowledge_base
            except Exception as e:
                logger.error(f"创建嵌入模型或知识库时出错: {str(e)}")
                raise ValueError(f"创建嵌入模型或知识库失败: {str(e)}. 请确保已安装sentence-transformers和faiss-cpu库")
        except Exception as e:
            logger.error(f"处理文本时出错: {str(e)}")
            raise

    def _assign_page_numbers_to_chunks(self, text: str, chunks: List[str], page_numbers: List[int]) -> Dict[str, int]:
        """
        为每个文本块分配页码

        参数:
            text: 完整文本
            chunks: 分割后的文本块
            page_numbers: 页码列表

        返回:
            page_info: 文本块到页码的映射字典
        """
        page_info = {}

        try:
            # 如果没有页码信息,则返回空字典
            if not page_numbers:
                logger.warning("没有页码信息,无法分配页码")
                return page_info

            # 简化页码分配算法,使用文本位置比例来估计页码
            text_length = len(text)

            for chunk in chunks:
                chunk_position = text.find(chunk)
                if chunk_position == -1:
                    # 如果找不到,使用最后一页
                    page_info[chunk] = page_numbers[-1]
                else:
                    # 根据文本位置比例估计页码
                    position_ratio = chunk_position / text_length
                    page_index = min(int(position_ratio * len(page_numbers)), len(page_numbers) - 1)
                    page_info[chunk] = page_numbers[page_index]

            return page_info
        except Exception as e:
            logger.error(f"分配页码时出错: {str(e)}")
            # 出错时返回空字典,但不中断程序
            return {}

    def answer_question(self, query: str, top_k: int = 4) -> Dict:
        """
        回答关于PDF内容的问题

        参数:
            query: 问题
            top_k: 检索的相关文档数量

        返回:
            result: 包含回答和来源页码的字典
        """
        if not self.knowledge_base:
            raise ValueError("知识库尚未初始化,请先调用load_pdf()")

        logger.info(f"处理问题: {query}")

        try:
            # 执行相似度搜索,找到与查询相关的文档
            docs = self.knowledge_base.similarity_search(query, k=top_k)

            llm = ChatOpenAI(
                openai_api_key= "sk-xxxxxx",  # 从环境变量获取 API Key
                openai_api_base="https://dashscope.aliyuncs.com/compatible-mode/v1",  # 百炼兼容端点
                model_name="qwen-plus",  # 指定阿里模型
                temperature=0.2  # 控制生成随机性
            )

            # llm = Ollama(
            #     model="deepseek-r1:1.5b",  # 直接指定模型名称
            #     base_url="http://localhost:11434",  # Ollama 服务地址
            #     temperature=0.2  # 控制随机性
            # )

            # 加载问答链
            chain = load_qa_chain(llm, chain_type="stuff")

            # 准备输入数据
            input_data = {"input_documents": docs, "question": query}

            # 使用回调函数跟踪API调用成本
            with get_openai_callback() as cost:
                # 执行问答链
                response = chain.invoke(input=input_data)
                logger.info(f"查询已处理。成本: {cost}")

            # 记录唯一的页码
            unique_pages = set()
            page_sources = []

            # 显示每个文档块的来源页码
            for i, doc in enumerate(docs):
                text_content = getattr(doc, "page_content", "")
                # 尝试直接匹配
                source_page = self.knowledge_base.page_info.get(text_content, None)

                # 如果直接匹配失败,尝试去除空白后匹配
                if source_page is None:
                    for chunk, page in self.knowledge_base.page_info.items():
                        if chunk.strip() == text_content.strip():
                            source_page = page
                            break

                # 如果仍然找不到匹配,使用默认值
                if source_page is None:
                    source_page = "未知"

                if source_page not in unique_pages:
                    unique_pages.add(source_page)
                    page_sources.append(source_page)

            # 按页码排序
            try:
                page_sources = sorted([p for p in page_sources if isinstance(p, int)])
            except Exception:
                # 如果排序失败,使用原始列表
                pass

            return {
                "answer": response["output_text"],
                "sources": page_sources,
                "cost": str(cost)
            }
        except Exception as e:
            logger.error(f"回答问题时出错: {str(e)}")
            raise

def main():
    """主函数,处理命令行参数并运行PDF聊天机器人"""
    parser = argparse.ArgumentParser(description="PDF聊天机器人 - 使用FAISS向量数据库回答PDF文档相关问题")
    parser.add_argument("--pdf", type=str, default="./xxx.pdf", help="PDF文件路径")
    parser.add_argument("--api_key", type=str, default="sk-xxxxxx", help="OpenAI API密钥")
    parser.add_argument("--base_url", type=str, default="http://localhost:11434/v1", help="API基础URL")
    parser.add_argument("--model_id", type=str, default="iic/nlp_gte_sentence-embedding_chinese-base", help="ModelScope模型ID")
    parser.add_argument("--device", type=str, help="计算设备,如 'cuda:0', 'cpu' 等")
    parser.add_argument("--no_query_prefix", action="store_true", help="不为查询添加前缀")
    parser.add_argument("--llm_model", type=str, default="deepseek-r1:1.5b", help="LLM模型名称")

    args = parser.parse_args()

    try:
        # 检查PDF文件是否存在
        if not os.path.exists(args.pdf):
            print(f"错误: PDF文件不存在: {args.pdf}")
            print("请提供有效的PDF文件路径")
            sys.exit(1)

        # 创建PDF聊天机器人
        chatbot = PDFChatBot(
            args.pdf, 
            args.api_key, 
            args.base_url, 
            args.model_id,
            args.device,
            not args.no_query_prefix,
            args.llm_model
        )

        # 加载PDF文件
        print(f"正在加载PDF文件: {args.pdf}...")
        try:
            chatbot.load_pdf()
        except Exception as e:
            print(f"加载PDF文件时出错: {str(e)}")
            print("请确保PDF文件格式正确且可读取")
            sys.exit(1)

        print(f"\n成功加载PDF文件: {args.pdf}")
        print(f"PDF共有 {chatbot.page_count} 页")
        print("\n现在您可以开始提问了! 输入'退出'或'exit'结束对话。\n")

        # 交互式问答循环
        while True:
            try:
                query = input("\n请输入您的问题: ")

                if query.lower() in ["退出", "exit", "quit", "q"]:
                    print("感谢使用PDF聊天机器人,再见!")
                    break

                if not query.strip():
                    continue

                try:
                    print("正在思考中...")
                    # 获取回答
                    result = chatbot.answer_question(query)

                    # 打印回答
                    print("\n回答:")
                    print(result["answer"])

                    # 打印来源
                    if result["sources"]:
                        print("\n来源页码:", ", ".join(map(str, result["sources"])))

                    # 打印成本
                    print(f"\n查询成本: {result['cost']}")

                except Exception as e:
                    print(f"处理问题时出错: {str(e)}")
                    print("请尝试重新提问或检查API密钥是否有效")
            except KeyboardInterrupt:
                print("\n程序被用户中断。感谢使用PDF聊天机器人,再见!")
                break
            except Exception as e:
                print(f"发生错误: {str(e)}")

    except Exception as e:
        print(f"程序运行出错: {str(e)}")

if __name__ == "__main__":
    try:
        main()
    except Exception as e:
        print(f"程序启动失败: {str(e)}")
        print("请检查环境配置和依赖库安装情况")

评论

发表回复