diff --git a/src/config/__init__.py b/src/config/__init__.py index 24d450a5..dbd9f8a5 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -49,7 +49,7 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths - self.add_item("model_provider", default="qianfan", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) + self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) self.add_item("model_name", default=None, des="模型名称") self.add_item("embed_model", default="bge-large-zh-v1.5", des="Embedding 模型", choices=["bge-large-zh-v1.5", "zhipu"]) self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"]) diff --git a/src/core/database.py b/src/core/database.py index 58751e40..c8dd43d8 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -2,10 +2,7 @@ import os import json import time from src.utils import hashstr, setup_logger, is_text_pdf -from src.plugins import pdf2txt -from src.core.knowledgebase import KnowledgeBase from src.core.filereader import pdfreader, plainreader -from src.core.graphbase import GraphDatabase from src.models.embedding import get_embedding_model logger = setup_logger("DataBaseManager") @@ -57,6 +54,8 @@ class DataBaseManager: self.embed_model = get_embedding_model(config) if self.config.enable_knowledge_base: + from src.core.knowledgebase import KnowledgeBase + from src.core.graphbase import GraphDatabase self.knowledge_base = KnowledgeBase(config, self.embed_model) self.graph_base = GraphDatabase(self.config, self.embed_model) self.graph_base.start() @@ -180,6 +179,7 @@ class DataBaseManager: if is_text_pdf(file): return pdfreader(file) else: + from src.plugins import pdf2txt return pdf2txt(file, return_text=True) elif file.endswith(".txt") or file.endswith(".md"): diff --git a/src/core/filereader.py b/src/core/filereader.py index ae9034dc..62fd11a6 100644 --- a/src/core/filereader.py +++ b/src/core/filereader.py @@ -1,14 +1,13 @@ import os from pathlib import Path -from llama_index.readers.file import PDFReader - def pdfreader(file_path): """读取PDF文件并返回text文本""" assert os.path.exists(file_path), "File not found" assert file_path.endswith(".pdf"), "File format not supported" + from llama_index.readers.file import PDFReader doc = PDFReader().load_data(file=Path(file_path)) # 简单的拼接起来之后返回纯文本 diff --git a/src/core/history.py b/src/core/history.py index 6b9fbced..98faa01c 100644 --- a/src/core/history.py +++ b/src/core/history.py @@ -25,9 +25,13 @@ class HistoryManager(): self.add_ai(content) return self.messages - def get_history_with_msg(self, msg, role="user"): + def get_history_with_msg(self, msg, role="user", max_rounds=None): """Get history with new message, but not append it to history.""" - history = self.messages[:] + if max_rounds is None: + history = self.messages[:] + else: + history = self.messages[-(2*max_rounds):] + history.append({"role": role, "content": msg}) return history diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 20eb68f0..7c0d5113 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -14,6 +14,7 @@ class KnowledgeBase: assert embed_model, "embed_model=None" self.embed_model = embed_model + self.client = MilvusClient(self.milvus_path) def _init_config(self, config): @@ -72,15 +73,17 @@ class KnowledgeBase: def search(self, query, collection_name, limit=3): query_vectors = self.embed_model.encode_queries([query]) + return self.search_by_vector(query_vectors[0], collection_name, limit) + def search_by_vector(self, vector, collection_name, limit=3): res = self.client.search( collection_name=collection_name, # target collection - data=query_vectors, # query vectors + data=[vector], # query vectors limit=limit, # number of returned entities output_fields=["text", "file_id"], # specifies fields to be returned ) - return res[0] # 因为 query 只有一个 + return res[0] def examples(self, collection_name, limit=20): res = self.client.query( diff --git a/src/core/retriever.py b/src/core/retriever.py index 19a93dff..6f74c7f4 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -15,9 +15,7 @@ class Retriever: def retrieval(self, query, history, meta): - refs = {} - refs["meta"] = meta - refs["rewritten_query"] = self.rewrite_query(query, history, refs) + refs = {"query": query, "history": history, "meta": meta} refs["entities"] = self.reco_entities(query, history, refs) refs["knowledge_base"] = self.query_knowledgebase(query, history, refs) refs["graph_base"] = self.query_graph(query, history, refs) @@ -68,37 +66,47 @@ class Retriever: def query_knowledgebase(self, query, history, refs): """查询知识库""" - query = refs.get("rewritten_query", query) + rw_query = self.rewrite_query(query, history, refs) kb_res = [] final_res = [] if refs["meta"].get("db_name") and self.config.enable_knowledge_base: + db_name = refs["meta"]["db_name"] - kb = self.dbm.metaname2db[refs["meta"]["db_name"]] + kb = self.dbm.metaname2db[db_name] limit = refs["meta"].get("queryCount", 10) - kb_res = self.dbm.knowledge_base.search(query, db_name, limit=limit) + + kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit) for r in kb_res: r["file"] = kb.id2file(r["entity"]["file_id"]) if self.config.enable_reranker: + RERANK_THRESHOLD = 0.1 for r in kb_res: r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) - final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1] + final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD] + else: final_res = kb_res[:5] - return {"results": final_res, "all_results": kb_res} + return {"results": final_res, "all_results": kb_res, "rw_query": rw_query} def rewrite_query(self, query, history, refs): """重写查询""" - if refs["meta"].get("rewrite_query") is None or history == []: + rewrite_query_span = refs["meta"].get("rewrite_query", None) + if rewrite_query_span is None or rewrite_query_span == "OFF": rewritten_query = query else: from src.utils.prompts import rewritten_query_prompt_template - rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query) + history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else "" + rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query) rewritten_query = self.model.predict(rewritten_query_prompt).content + if rewrite_query_span == "HyDE": + hy_doc = self.model.predict(rewritten_query).content + rewritten_query = f"{rewritten_query} {hy_doc}" + return rewritten_query def reco_entities(self, query, history, refs): diff --git a/src/models/README.md b/src/models/README.md index 27e344ac..433e24ce 100644 --- a/src/models/README.md +++ b/src/models/README.md @@ -2,12 +2,12 @@ ### 1. 对话模型支持 -模型仅支持通过API调用的模型,如果是需要运行本地模型,则建议使用 vllm 转成 API 服务之后使用。 +模型仅支持通过API调用的模型,如果是需要运行本地模型,则建议使用 vllm 转成 API 服务之后使用。使用前请配置 APIKEY 后使用,配置项目参考:[.env.template](../.env.template) |模型供应商(`config.model_provider`)|默认模型(`config.model_name`)|配置项目(`.env`)| |:-|:-|:-| |`qianfan`|`ernie_speed`|`QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY`| -|`zhipu`|`glm-4`|`ZHIPUAPI`| +|`zhipu`(default)|`glm-4`|`ZHIPUAPI`| |`deepseek`|`deepseek-chat`|`DEEPSEEKAPI`| |`vllm`|`vllm`|`VLLM_API_KEY`, `VLLM_API_BASE`| diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 286fe5d0..da3c7765 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -64,7 +64,6 @@ class VLLM(OpenAIBase): super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) -import qianfan class GeneralResponse: @@ -76,6 +75,7 @@ class GeneralResponse: class Qianfan: def __init__(self, model_name="ernie_speed") -> None: + import qianfan self.model_name = model_name access_key = os.getenv("QIANFAN_ACCESS_KEY") secret_key = os.getenv("QIANFAN_SECRET_KEY") diff --git a/src/plugins/pdf2txt.py b/src/plugins/pdf2txt.py index e823d021..a0e4f255 100644 --- a/src/plugins/pdf2txt.py +++ b/src/plugins/pdf2txt.py @@ -1,38 +1,19 @@ import os import fitz # fitz就是pip install PyMuPDF import cv2 -from paddleocr import PPStructure, save_structure_res -from paddleocr.ppstructure.recovery.recovery_to_doc import sorted_layout_boxes, convert_info_docx from copy import deepcopy from tqdm import tqdm def pdf2txt(pdf_path, return_text=False): + from paddleocr import PPStructure, save_structure_res + from paddleocr.ppstructure.recovery.recovery_to_doc import sorted_layout_boxes, convert_info_docx output_dir = os.path.join('tmp', 'pdf2txt', os.path.basename(pdf_path).split('.')[0]) os.makedirs(output_dir, exist_ok=True) table_engine = PPStructure(recovery=True, lang='ch') - imgs = [] - img_dir = os.path.join(output_dir, 'imgs') - if not os.path.exists(img_dir): - os.makedirs(img_dir) - pdfDoc = fitz.open(pdf_path) - totalPage = pdfDoc.page_count - for pg in tqdm(range(totalPage), desc='to imgs', ncols=100): - page = pdfDoc[pg] - rotate = int(0) - zoom_x = 2 - zoom_y = 2 - mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate) - pix = page.get_pixmap(matrix=mat, alpha=False) - img_filename = os.path.join(img_dir, f'images_{pg+1}.png') - pix.save(img_filename) # os.sep - imgs.append(img_filename) - else: - img_names = sorted(os.listdir(img_dir)) - imgs = [os.path.join(img_dir, img_name) for img_name in img_names] - + imgs = convert_imgs(pdf_path, output_dir) respath = os.path.join(output_dir, 'res.txt') text = [] @@ -59,7 +40,7 @@ def pdf2txt(pdf_path, return_text=False): continue # 如果不是字典或者缺少 'text' 键,跳过当前循环 text.append('\n') - whole_text = '\n'.join(text) + whole_text = ''.join(text) with open(respath, 'w', encoding='utf-8') as f: f.write(whole_text) @@ -68,6 +49,29 @@ def pdf2txt(pdf_path, return_text=False): return respath +def convert_imgs(pdf_path, output_dir): + imgs = [] + img_dir = os.path.join(output_dir, 'imgs') + if not os.path.exists(img_dir): + os.makedirs(img_dir) + pdfDoc = fitz.open(pdf_path) + totalPage = pdfDoc.page_count + for pg in tqdm(range(totalPage), desc='to imgs', ncols=100): + page = pdfDoc[pg] + rotate = int(0) + zoom_x = 2 + zoom_y = 2 + mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate) + pix = page.get_pixmap(matrix=mat, alpha=False) + img_filename = os.path.join(img_dir, f'images_{pg+1}.png') + pix.save(img_filename) # os.sep + imgs.append(img_filename) + else: + img_names = sorted(os.listdir(img_dir)) + imgs = [os.path.join(img_dir, img_name) for img_name in img_names] + + return imgs + if __name__ == "__main__": - pdf_path = r'data/file/焙烤食品工艺学.pdf' + pdf_path = r'saves/data/uploads/2e04d5_保健食品.pdf' print(pdf2txt(pdf_path)) diff --git a/src/requirements.txt b/src/requirements.txt index 32f414e5..29577253 100644 --- a/src/requirements.txt +++ b/src/requirements.txt @@ -1,9 +1,6 @@ FlagEmbedding==1.2.10 Flask==3.0.3 Flask_Cors==4.0.1 -llama_index==0.10.53 openai==1.35.10 -pymilvus==2.4.4 python-dotenv==1.0.1 PyYAML==6.0.1 -qianfan==0.4.0.1 diff --git a/src/views/common_view.py b/src/views/common_view.py index 221264ad..d9445401 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -34,7 +34,7 @@ def chat(): new_query, refs = startup.retriever(query, history_manager.messages, meta) - messages = history_manager.get_history_with_msg(new_query) + messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round')) history_manager.add_user(query) logger.debug(f"Web history: {history_manager}") diff --git a/web/src/assets/base.css b/web/src/assets/base.css index 3d4e5f0e..6e5c9c8d 100644 --- a/web/src/assets/base.css +++ b/web/src/assets/base.css @@ -45,7 +45,6 @@ --main-light-6: #FAFDFD; --min-width: 400px; --min-header-width: 80px; - --min-sider-width: 100px; --error-color: #f50a0d; } diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index b8a99065..5861a5e1 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -39,11 +39,11 @@ -