diff --git a/.gitignore b/.gitignore index 409b49b2..f602ffa0 100644 --- a/.gitignore +++ b/.gitignore @@ -40,4 +40,5 @@ local_neo4j/data local_neo4j/logs local_neo4j/import local_neo4j/plugins -local_neo4j/conf \ No newline at end of file +local_neo4j/conf +graphrag \ No newline at end of file diff --git a/local_neo4j/test_neo4j.py b/local_neo4j/test_neo4j.py index a293e4fc..53b967a0 100644 --- a/local_neo4j/test_neo4j.py +++ b/local_neo4j/test_neo4j.py @@ -1,15 +1,17 @@ from neo4j import GraphDatabase from neo4j.exceptions import ServiceUnavailable, AuthError +# sudo ln -s /snap/core22/1586/usr/sbin/iptables /usr/sbin/iptables + def check_neo4j_status(uri="bolt://localhost:7687", username="neo4j", password="0123456789"): """ 检查 Neo4j 数据库是否可以连接并正常工作。 - + 参数: uri (str): Neo4j 的 URI,默认为 "bolt://localhost:7687" username (str): 数据库用户名,默认为 "neo4j" password (str): 数据库密码,默认为 "0123456789" - + 返回: str: "OK" 表示连接成功,"UNAVAILABLE" 表示服务不可用,"AUTH_FAILED" 表示认证失败。 """ diff --git a/src/config/__init__.py b/src/config/__init__.py index 50ae6ce7..1828e406 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="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) + self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["openai", "qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) self.add_item("model_name", default=None, des="模型名称") self.add_item("embed_model", default="zhipu-embedding-3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"]) @@ -145,6 +145,13 @@ class Config(SimpleConfig): MODEL_NAMES = { # https://platform.deepseek.com/api-docs/zh-cn/pricing + "openai": [ + "gpt-4", + "gpt-4o", + "gpt-4-0125-preview", + "gpt-4o-mini", + ], + "deepseek": [ "deepseek-chat", "deepseek-coder" diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 664b4729..f6db98af 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -15,6 +15,71 @@ logger = setup_logger("server-graphbase") warnings.filterwarnings("ignore", category=UserWarning) + +""" +from neo4j import GraphDatabase +import random + +class KnowledgeGraph: + def __init__(self, uri, user, password): + self._driver = GraphDatabase.driver(uri, auth=(user, password)) + + def close(self): + self._driver.close() + + def use_database(self, kgdb_name): + with self._driver.session() as session: + session.run(f"USE {kgdb_name}") + + def get_sample_nodes(self, kgdb_name='neo4j', num=50): + self.use_database(kgdb_name) + selected_nodes = set() + nodes_to_expand = set() + result_nodes = [] + + with self._driver.session() as session: + while len(selected_nodes) < num: + # 如果需要扩展的节点为空,随机选择一个新节点 + if not nodes_to_expand: + result = session.run("MATCH (n) RETURN n, rand() as r ORDER BY r LIMIT 1") + for record in result: + nodes_to_expand.add(record['n'].id) + result_nodes.append({'n': record['n'], 'r': None, 'm': None}) + + # 从需要扩展的节点中随机选择一个节点 + current_node_id = random.choice(list(nodes_to_expand)) + nodes_to_expand.remove(current_node_id) + + # 获取当前节点的邻居节点,最多5个 + result = session.run( + f"MATCH (n)-[r]-(m) WHERE id(n) = {current_node_id} RETURN n, r, m LIMIT 5" + ) + + for record in result: + neighbor_node_id = record['m'].id + if neighbor_node_id not in selected_nodes: + selected_nodes.add(neighbor_node_id) + nodes_to_expand.add(neighbor_node_id) + result_nodes.append({'n': record['n'], 'r': record['r'], 'm': record['m']}) + + # 如果已经达到最大值,停止扩展 + if len(selected_nodes) >= num: + break + + return result_nodes[:num] + +# 示例用法 +uri = "bolt://localhost:7687" +user = "neo4j" +password = "password" +kg = KnowledgeGraph(uri, user, password) +sample_nodes = kg.get_sample_nodes(num=50) +for node in sample_nodes: + print(f"Node: {node['n'].id}, Relationship: {node['r']}, Neighbor: {node['m'].id}") +kg.close() + +""" + UIE_MODEL = None class GraphDatabase: @@ -38,7 +103,7 @@ class GraphDatabase: self.driver.close() def get_sample_nodes(self, kgdb_name='neo4j', num=50): - """获取指定数据库的前 num 个节点信息""" + """获取指定数据库的 num 个节点信息""" self.use_database(kgdb_name) def query(tx, num): result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num)) diff --git a/src/core/history.py b/src/core/history.py index 98faa01c..dae392e6 100644 --- a/src/core/history.py +++ b/src/core/history.py @@ -38,5 +38,6 @@ class HistoryManager(): def __str__(self): history_str = "" for message in self.messages: - history_str += f"{message['role']}: {message['content']}" + msg = message["content"].replace('\n', ' ') + history_str += f"\n{message['role']}: {msg}" return history_str \ No newline at end of file diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 327384f0..873f925a 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -63,7 +63,7 @@ class KnowledgeBase: "id": int(random.random() * 1e12), "vector": vectors[i], "text": docs[i], - "hash": hashstr(docs[i] + str(random.random())), + "hash": hashstr(docs[i], with_salt=True), **kwargs} for i in range(len(vectors))] res = self.client.insert(collection_name=collection_name, data=data) diff --git a/src/core/retriever.py b/src/core/retriever.py index bafa7f84..d2bcd128 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -84,28 +84,33 @@ class Retriever: db_name = refs["meta"]["db_name"] kb = self.dbm.metaname2db[db_name] - limit = refs["meta"].get("queryCount", 10) - kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit) - for r in kb_res: + max_query_count = refs["meta"].get("maxQueryCount", 10) + rerank_threshold = refs["meta"].get("rerankThreshold", 0.1) + distance_threshold = refs["meta"].get("distanceThreshold", 0) + top_k = refs["meta"].get("topK", 5) + + all_kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count) + for r in all_kb_res: r["file"] = kb.id2file(r["entity"]["file_id"]) + # use distance threshold to filter results + kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] + if self.config.enable_reranker: - RERANK_THRESHOLD = 0.001 for r in kb_res: r["rerank_score"] = self.reranker.compute_score([rw_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"] > RERANK_THRESHOLD] + kb_res = [_res for _res in kb_res if _res["rerank_score"] > rerank_threshold] - else: - final_res = kb_res[:5] + kb_res = kb_res[:top_k] - return {"results": final_res, "all_results": kb_res, "rw_query": rw_query} + return {"results": kb_res, "all_results": all_kb_res, "rw_query": rw_query} def rewrite_query(self, query, history, refs): """重写查询""" - rewrite_query_span = refs["meta"].get("rewrite_query", None) - if rewrite_query_span is None or rewrite_query_span == "OFF": + rewrite_query_span = refs["meta"].get("rewriteQuery", "off") + if rewrite_query_span == "off": rewritten_query = query else: from src.utils.prompts import rewritten_query_prompt_template @@ -114,7 +119,7 @@ class Retriever: 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": + if rewrite_query_span == "hyde": hy_doc = self.model.predict(rewritten_query).content rewritten_query = f"{rewritten_query} {hy_doc}" @@ -168,7 +173,6 @@ class Retriever: return node_info, edge_info def format_general_results(self, results): - logger.debug(f"Formatting general results: {results}") formatted_results = {"nodes": [], "edges": []} for item in results: @@ -189,7 +193,6 @@ class Retriever: return formatted_results def format_query_results(self, results): - logger.debug(f"Formatting query results: {results}") formatted_results = {"nodes": [], "edges": []} node_dict = {} diff --git a/src/models/README.md b/src/models/README.md index 433e24ce..fe1fc347 100644 --- a/src/models/README.md +++ b/src/models/README.md @@ -38,7 +38,7 @@ python -m vllm.entrypoints.openai.api_server \ |`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`(*修改为本地路径)| |`zhipu`|`embedding-2`|`ZHIPUAPI` (`.env`)| -### 3. 重排序模型支持 + 例如(`config/base.yaml`): @@ -52,3 +52,7 @@ model_name: null # for default model_local_paths: bge-large-zh-v1.5: /models/bge-large-zh-v1.5 ``` + +### 3. 重排序模型支持 + +目前仅支持 `BAAI/bge-reranker-v2-m3`。 diff --git a/src/models/__init__.py b/src/models/__init__.py index 857e0b80..8e1f6592 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -28,6 +28,10 @@ def select_model(config): from src.models.chat_model import DashScope return DashScope(model_name) + elif model_provider == "openai": + from src.models.chat_model import OpenModel + return OpenModel(model_name) + elif model_provider is None: raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`") else: diff --git a/src/models/chat_model.py b/src/models/chat_model.py index be612215..2ed63c18 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -39,6 +39,14 @@ class OpenAIBase(): return response.choices[0].message +class OpenModel(OpenAIBase): + def __init__(self, model_name=None): + model_name = model_name or "gpt-4o-mini" + api_key = os.getenv("OPENAI_API_KEY") + base_url = os.getenv("OPENAI_API_BASE") + super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) + + class DeepSeek(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "deepseek-chat" diff --git a/src/models/dify_fording.py b/src/models/dify_fording.py deleted file mode 100644 index e69de29b..00000000 diff --git a/src/models/embedding.py b/src/models/embedding.py index 9eb81ee2..21d32a03 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,12 +1,16 @@ import os +import uuid from FlagEmbedding import FlagModel, FlagReranker from src.config import EMBED_MODEL_INFO, RERANKER_LIST from src.utils.logging_config import setup_logger +from src.utils import hashstr logger = setup_logger("EmbeddingModel") +GLOBAL_EMBED_STATE = {} + class EmbeddingModel(FlagModel): def __init__(self, model_info, config, **kwargs): @@ -47,9 +51,21 @@ class ZhipuEmbedding: def predict(self, message): data = [] + if len(message) > 10: + global GLOBAL_EMBED_STATE + task_id = hashstr(message) + logger.info(f"Creating new state for process {task_id}") + GLOBAL_EMBED_STATE[task_id] = { + 'status': 'in-progress', + 'total': len(message), + 'progress': 0 + } + for i in range(0, len(message), 10): if len(message) > 10: logger.info(f"Encoding {i} to {i+10} with {len(message)} messages") + GLOBAL_EMBED_STATE[task_id]['progress'] = i + group_msg = message[i:i+10] response = self.client.embeddings.create( model=self.model_info.default_path, @@ -58,6 +74,10 @@ class ZhipuEmbedding: data.extend([a.embedding for a in response.data]) + if len(message) > 10: + GLOBAL_EMBED_STATE[task_id]['progress'] = len(message) + GLOBAL_EMBED_STATE[task_id]['status'] = 'completed' + return data def encode(self, message): diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index 9f150d26..aea4c5c8 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -16,7 +16,7 @@ logger = setup_logger("OneKE") dotenv.load_dotenv() -MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR'), 'OneKE') +MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR', './'), 'OneKE') instruction_mapper = { 'NERzh': "你是专门进行实体抽取的专家。请从input中抽取出符合schema定义的实体,不存在的实体类型返回空列表。请按照JSON字符串的格式回答。", diff --git a/src/plugins/pdf2txt.py b/src/plugins/pdf2txt.py index a0e4f255..ae64ea39 100644 --- a/src/plugins/pdf2txt.py +++ b/src/plugins/pdf2txt.py @@ -1,25 +1,53 @@ import os +import uuid import fitz # fitz就是pip install PyMuPDF import cv2 from copy import deepcopy from tqdm import tqdm +from argparse import ArgumentParser +from src.utils import logger + +GOLBAL_STATE = {} def pdf2txt(pdf_path, return_text=False): + if not os.path.exists(pdf_path): + raise FileNotFoundError(f"File not found: {pdf_path}") + + # Importing these modules here to avoid unnecessary imports in other files 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]) + filename = os.path.basename(pdf_path).split('.')[0] + output_dir = os.path.join('saves', 'data', 'pdf2txt', filename) os.makedirs(output_dir, exist_ok=True) table_engine = PPStructure(recovery=True, lang='ch') imgs = convert_imgs(pdf_path, output_dir) - respath = os.path.join(output_dir, 'res.txt') + respath = os.path.join(output_dir, f'{filename}.txt') + + global GOLBAL_STATE + task_id = str(uuid.uuid4()) + if task_id in GOLBAL_STATE: + logger.info(f"Reusing previous state for process {task_id}") + return GOLBAL_STATE[task_id] + else: + logger.info(f"Creating new state for process {task_id}") + GOLBAL_STATE[task_id] = { + 'pdf_path': pdf_path, + 'return_text': return_text, + 'output_dir': output_dir, + 'status': 'in-progress', + 'total': len(imgs), + 'progress': 0 + } + text = [] for img_name in tqdm(imgs, desc='to txt', ncols=100): img = cv2.imread(img_name) result = table_engine(img) + GOLBAL_STATE[task_id]['progress'] += 1 save_structure_res(result, output_dir, "structure_result") @@ -43,6 +71,9 @@ def pdf2txt(pdf_path, return_text=False): whole_text = ''.join(text) with open(respath, 'w', encoding='utf-8') as f: f.write(whole_text) + logger.info(f"Extracted text saved to {respath}") + + GOLBAL_STATE[task_id]['status'] = 'completed' if return_text: return whole_text @@ -72,6 +103,13 @@ def convert_imgs(pdf_path, output_dir): return imgs +def get_state(task_id): + return GOLBAL_STATE.get(task_id, {}) + if __name__ == "__main__": - pdf_path = r'saves/data/uploads/2e04d5_保健食品.pdf' - print(pdf2txt(pdf_path)) + parser = ArgumentParser() + parser.add_argument('--pdf-path', type=str, required=True, help='Path to the PDF file') + parser.add_argument('--return-text', action='store_true', help='Return the extracted text') + args = parser.parse_args() + + pdf2txt(args.pdf_path, args.return_text) diff --git a/src/utils/__init__.py b/src/utils/__init__.py index 1df9e7e6..5b613e7a 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -1,4 +1,5 @@ import time +import random from src.utils.logging_config import setup_logger, logger def is_text_pdf(pdf_path): @@ -15,7 +16,7 @@ def hashstr(input_string, length=8, with_salt=False): import hashlib # 添加时间戳作为干扰 if with_salt: - input_string += str(time.time()) + input_string += str(time.time() + random.random()) hash = hashlib.md5(str(input_string).encode()).hexdigest() return hash[:length] \ No newline at end of file diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index 2eb64b50..67a4d1f9 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -5,10 +5,10 @@ from datetime import datetime DATETIME = datetime.now().strftime('%Y-%m-%d-%H%M%S') # DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件 -LOG_FILE = f'log/project-{DATETIME}.log' +LOG_FILE = f'saves/log/project-{DATETIME}.log' def setup_logger(name, level=logging.DEBUG, console=False): - os.makedirs("log", exist_ok=True) + os.makedirs("saves/log", exist_ok=True) """Function to setup logger with the given name and log file.""" logger = logging.getLogger(name) diff --git a/src/views/common_view.py b/src/views/common_view.py index 3eec200f..730ee785 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -1,3 +1,4 @@ +import itertools import json from flask import Blueprint, jsonify, request, Response @@ -89,6 +90,7 @@ def restart(): def get_log(): from src.utils.logging_config import LOG_FILE with open(LOG_FILE, 'r') as f: - log = f.read() + log_lines = itertools.islice(f, 1000) + log = ''.join(log_lines) return jsonify({"log": log}) \ No newline at end of file diff --git a/web/public/icon-narrow.png b/web/public/icon-narrow.png new file mode 100644 index 00000000..5fa011b2 Binary files /dev/null and b/web/public/icon-narrow.png differ diff --git a/web/public/icon.png b/web/public/icon.png new file mode 100644 index 00000000..8b200a98 Binary files /dev/null and b/web/public/icon.png differ diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 3b72d028..dd602682 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -30,10 +30,10 @@