diff --git a/README.md b/README.md index 208feb00..b33c2155 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,6 @@ - ## 概述 基于大模型 RAG 知识库与知识图谱的问答平台。Llamaindex + VueJS + Flask + Neo4j。大模型适配 OpenAI、国内主流大模型平台的模型调用、本地 vllm 部署。只需要配置对应服务平台的 `API_KEY` 即可使用。 @@ -18,11 +17,12 @@ ![image](https://github.com/user-attachments/assets/75010511-4ac5-4924-8268-fea9a589839c) -代办清单 +待办清单 - [ ] Ollma Embedding 支持(Open-like Embedding 支持) -- [x] 知识图谱索引支持自定义 Embedding 模型 -- [x] DeepSeek-R1 支持 +- [ ] embedding 过程添加进度显示 +- [X] 知识图谱索引支持自定义 Embedding 模型 +- [X] DeepSeek-R1 支持 ## 更新日志 @@ -39,7 +39,6 @@ - ## 快速上手 在启动之前,提供 API 服务商的 API_KEY,并放置在 `src/.env` 文件中。默认使用的是硅基流动。因此务必需要配置 `SILICONFLOW_API_KEY=`。其余模型的配置可以参考 [src/static/models.yaml](src/static/models.yaml) 中的 env。 @@ -103,15 +102,15 @@ docker compose -f docker/docker-compose.yml --env-file src/.env up --build 模型仅支持通过API调用的模型,如果是需要运行本地模型,则建议使用 vllm 转成 API 服务之后使用。使用前请在 `.env` 配置 APIKEY 后使用,配置项目参考:[src/static/models.yaml](src/static/models.yaml) -| 模型供应商 | 默认模型 | 配置项目 | -| :-------------------- | :---------------------------------------- | :--------------------------------------------- | -| `siliconflow` (默认) | `Qwen/Qwen2.5-7B-Instruct` (免费) | `SILICONFLOW_API_KEY` | -| `openai` | `gpt-4o` | `OPENAI_API_KEY` | -| `deepseek` | `deepseek-chat` | `DEEPSEEK_API_KEY` | -| `arc`(豆包方舟) | `doubao-1-5-pro-32k-250115` | `ARK_API_KEY` | -| `zhipu`(智谱清言) | `glm-4-flash` | `ZHIPUAI_API_KEY` | -| `dashscope`(阿里) | `qwen-max-latest` | `DASHSCOPE_API_KEY` | -| `qianfan`(百度) | `ernie_speed` | `QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY` | +| 模型供应商 | 默认模型 | 配置项目 | +| :--------------------- | :---------------------------------- | :--------------------------------------------- | +| `siliconflow` (默认) | `Qwen/Qwen2.5-7B-Instruct` (免费) | `SILICONFLOW_API_KEY` | +| `openai` | `gpt-4o` | `OPENAI_API_KEY` | +| `deepseek` | `deepseek-chat` | `DEEPSEEK_API_KEY` | +| `arc`(豆包方舟) | `doubao-1-5-pro-32k-250115` | `ARK_API_KEY` | +| `zhipu`(智谱清言) | `glm-4-flash` | `ZHIPUAI_API_KEY` | +| `dashscope`(阿里) | `qwen-max-latest` | `DASHSCOPE_API_KEY` | +| `qianfan`(百度) | `ernie_speed` | `QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY` | 此外,如果想要添加供应商的模型,确认知识 OpenAI 调用的方法之后,只需要在 [src/static/models.yaml](src/static/models.yaml) 中添加对应的模型配置即可。配置示例如下: @@ -140,14 +139,12 @@ docker compose -f docker/docker-compose.yml --env-file src/.env up --build > [!Warning] > 需要注意,由于知识库和图数据库的构建都依赖于向量模型,如果中途更改向量模型,会导致知识库不可用。此外,知识图谱的向量索引的建立默认使用 embedding-3 构建,因此检索的时候必须使用 embedding-3(现阶段还不支持修改) - 对于**语言模型**,并不支持直接运行本地语言模型,请使用 vllm 或者 ollama 转成 API 服务之后使用。 对于**向量模型**和**重排序模型**,选择以 `local` 前缀开头的模型,可以不做修改会自动下载模型,如果下载过程中出现问题,请参考 [HF-Mirror](https://hf-mirror.com/) 配置相关内容。(但请注意,如果是 Docker 运行,模型仅会缓存到 Docker 里面) 如果想要使用本地已经下载好的模型,可以在网页的 settings 里面做映射。或者修改 `saves/config/base.yaml` 来配置映射关系。但请记得,本地模型的路径要在 docker-compose 的文件中映射 volumes。 - ## 知识库支持 本项目支持多种格式的知识库,目前支持的格式有:pdf、txt、md。文件上传后,会首先转换为纯文本,然后使用向量模型将文本转换为向量,然后存储到向量数据库中。这个过程可能会花费不少时间。 diff --git a/docker/web.Dockerfile b/docker/web.Dockerfile index 20737eb3..3a3624fb 100644 --- a/docker/web.Dockerfile +++ b/docker/web.Dockerfile @@ -6,8 +6,8 @@ WORKDIR /app COPY ./web/package*.json ./ # 安装依赖 -RUN npm install --verbose --force -# RUN npm install --registry http://mirrors.cloud.tencent.com/npm/ --verbose --force +# RUN npm install --verbose --force +RUN npm install --registry http://mirrors.cloud.tencent.com/npm/ --verbose --force # 复制源代码 COPY ./web . @@ -22,8 +22,8 @@ FROM node:latest AS build-stage WORKDIR /app COPY ./web/package*.json ./ -RUN npm install --force -# RUN npm install --registry https://registry.npmmirror.com --force +# RUN npm install --force +RUN npm install --registry https://registry.npmmirror.com --force COPY ./web . RUN npm run build diff --git a/src/config/__init__.py b/src/config/__init__.py index 32a361ea..2de995f0 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -2,11 +2,9 @@ import os import json import yaml from pathlib import Path -from src.utils.logging_config import setup_logger +from src.utils.logging_config import logger -logger = setup_logger("Config") - -with open(Path("src/static/models.yaml"), "r") as f: +with open(Path("src/static/models.yaml"), 'r', encoding='utf-8') as f: _models = yaml.safe_load(f) MODEL_NAMES = _models["MODEL_NAMES"] @@ -53,7 +51,6 @@ class Config(SimpleConfig): self.add_item("enable_reranker", default=False, des="是否开启重排序") self.add_item("enable_knowledge_base", default=False, des="是否开启知识库") self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱") - self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎") self.add_item("enable_web_search", default=False, des="是否开启网页搜索(需配置 TAVILY_API_KEY)") # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 @@ -90,6 +87,9 @@ class Config(SimpleConfig): def handle_self(self): self.model_names = MODEL_NAMES + self.embed_model_names = EMBED_MODEL_INFO + self.reranker_names = RERANKER_LIST + model_provider_info = self.model_names.get(self.model_provider, {}) self.model_dir = os.environ.get("MODEL_DIR", "") diff --git a/src/core/database.py b/src/core/database.py index 3498dbb8..2b80e22c 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -2,11 +2,9 @@ import os import json import time from src.plugins import pdf2txt -from src.utils import hashstr, setup_logger, is_text_pdf +from src.utils import hashstr, logger, is_text_pdf from src.models.embedding import get_embedding_model -logger = setup_logger("DataBaseManager") - class DataBaseManager: @@ -21,7 +19,6 @@ class DataBaseManager: if self.config.enable_knowledge_graph: from src.core.graphbase import GraphDatabase self.graph_base = GraphDatabase(self.config, self.embed_model) - self.graph_base.start() else: self.graph_base = None diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 399ec448..6294c683 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -9,9 +9,7 @@ import warnings from src.plugins import pdf2txt from src.plugins.oneke import OneKE -from src.utils import setup_logger - -logger = setup_logger("server-graphbase") +from src.utils import logger warnings.filterwarnings("ignore", category=UserWarning) @@ -28,6 +26,16 @@ class GraphDatabase: assert embed_model, "embed_model=None" self.embed_model = embed_model self.embed_model_name = None + self.work_dir = os.path.join(config.save_dir, "knowledge_graph", kgdb_name) + os.makedirs(self.work_dir, exist_ok=True) + + # 尝试加载已保存的图数据库信息 + if self.load_graph_info(): + logger.info(f"已加载图数据库信息:{self.kgdb_name}") + else: + logger.info(f"未找到已保存的图数据库信息,将创建新的配置") + + self.start() def start(self): uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687") @@ -38,11 +46,12 @@ class GraphDatabase: self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password)) self.status = "open" logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_database_info()}") + # 连接成功后保存图数据库信息 + self.save_graph_info() except Exception as e: logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}") self.config.enable_knowledge_graph = False - def close(self): """关闭数据库连接""" self.driver.close() @@ -119,33 +128,39 @@ class GraphDatabase: with self.driver.session() as session: session.execute_write(create, triples) - def pdf_file_add_entity(self, file_path, output_path, kgdb_name='neo4j'): - self.use_database(kgdb_name) # 切换到指定数据库 - text_path = pdf2txt(file_path) - global UIE_MODEL - if UIE_MODEL is None: - UIE_MODEL = OneKE() - triples_path = UIE_MODEL.processing_text_to_kg(text_path, output_path) - self.jsonl_file_add_entity(triples_path) - return kgdb_name + # def pdf_file_add_entity(self, file_path, output_path, kgdb_name='neo4j'): + # self.use_database(kgdb_name) # 切换到指定数据库 + # text_path = pdf2txt(file_path) + # global UIE_MODEL + # if UIE_MODEL is None: + # UIE_MODEL = OneKE() + # triples_path = UIE_MODEL.processing_text_to_kg(text_path, output_path) + # self.jsonl_file_add_entity(triples_path) + # return kgdb_name def txt_add_vector_entity(self, triples, kgdb_name='neo4j'): """添加实体三元组""" self.use_database(kgdb_name) def _index_exists(tx, index_name): + """检查索引是否存在""" result = tx.run("SHOW INDEXES") for record in result: if record["name"] == index_name: return True return False + def _create_graph(tx, data): + """添加一个三元组""" for entry in data: tx.run(""" MERGE (h:Entity {name: $h}) MERGE (t:Entity {name: $t}) MERGE (h)-[r:RELATION {type: $r}]->(t) """, h=entry['h'], t=entry['t'], r=entry['r']) + def _create_vector_index(tx, dim): + """创建向量索引""" + # NOTE 这里是否是会重复构建索引? index_name = "entityEmbeddings" if not _index_exists(tx, index_name): tx.run(f""" @@ -157,24 +172,34 @@ class GraphDatabase: }} }}; """) + # 判断模型名称是否匹配 from src.config import EMBED_MODEL_INFO - embed_info = EMBED_MODEL_INFO[self.config.embed_model] - with self.driver.session() as session: - session.execute_write(_create_graph, triples) - session.execute_write(_create_vector_index, embed_info.get('dimension')) - for i, entry in enumerate(triples): - logger.info(f"Adding entity {i+1}/{len(triples)}") - embedding_h = self.get_embedding(entry['h']) - session.execute_write(self.set_embedding, entry['h'], embedding_h) + cur_embed_info = EMBED_MODEL_INFO[self.config.embed_model] + self.embed_model_name = self.embed_model_name or cur_embed_info.get('name') + assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \ + f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}" + with self.driver.session() as session: + logger.info(f"Adding entity to {kgdb_name}") + session.execute_write(_create_graph, triples) + logger.info(f"Creating vector index for {kgdb_name} with {self.config.embed_model}") + session.execute_write(_create_vector_index, cur_embed_info['dimension']) + # NOTE 这里需要异步处理 + for i, entry in enumerate(triples): + logger.debug(f"Adding entity {i+1}/{len(triples)}") + embedding_h = self.get_embedding(entry['h']) embedding_t = self.get_embedding(entry['t']) + session.execute_write(self.set_embedding, entry['h'], embedding_h) session.execute_write(self.set_embedding, entry['t'], embedding_t) + # 数据添加完成后保存图信息 + self.save_graph_info() + def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'): self.status = "processing" kgdb_name = kgdb_name or 'neo4j' - self.embed_model_name = self.embed_model_name or self.config.embed_model self.use_database(kgdb_name) # 切换到指定数据库 + logger.info(f"Start adding entity to {kgdb_name} with {file_path}") def read_triples(file_path): with open(file_path, 'r', encoding='utf-8') as file: @@ -186,6 +211,8 @@ class GraphDatabase: self.txt_add_vector_entity(triples, kgdb_name) self.status = "open" + # 更新并保存图数据库信息 + self.save_graph_info() return kgdb_name def delete_entity(self, entity_name=None, kgdb_name="neo4j"): @@ -213,27 +240,14 @@ class GraphDatabase: def query_node(self, entity_name, hops=2, **kwargs): # TODO 添加判断节点数量为 0 停止检索 + + logger.debug(f"Query graph node {entity_name} with {hops=}") if kwargs.get("exact_match"): raise NotImplemented("not implement for `exact_match`") else: return self.query_by_vector(entity_name=entity_name, **kwargs) def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=5): - results = self.query_by_vector_tep(entity_name=entity_name) - - # 筛选出分数高于阈值的实体 - qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] - - # 对每个合格的实体进行查询 - all_query_results = [] - for entity in qualified_entities: - query_result = self.query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name) - all_query_results.extend(query_result) - - return all_query_results - - def query_by_vector_tep(self, entity_name, kgdb_name='neo4j'): - """向量查询""" self.use_database(kgdb_name) def query(tx, text): embedding = self.get_embedding(text) @@ -245,15 +259,28 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - return session.execute_read(query, entity_name) + results = session.execute_read(query, entity_name) + + + # 筛选出分数高于阈值的实体 + qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] + logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") + + # 对每个合格的实体进行查询 + all_query_results = [] + for entity in qualified_entities: + query_result = self.query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name) + all_query_results.extend(query_result) + + return all_query_results def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2): - """查询指定实体三元组信息""" + """查询指定实体三元组信息(无向关系)""" self.use_database(kgdb_name) def query(tx, entity_name, hops): result = tx.run(f""" - MATCH (n {{name: $entity_name}})-[r*1..{hops}]->(m) - RETURN n.name AS node_name, r, m.name AS neighbor_name + MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m) + RETURN n, r, m """, entity_name=entity_name) return result.values() @@ -316,11 +343,9 @@ class GraphDatabase: return session.execute_read(query, node_name, hops) def get_embedding(self, text): - inputs = [text] with torch.no_grad(): - outputs = self.embed_model.encode(inputs) - embeddings = outputs[0] # 假设取平均作为文本的嵌入向量 - return embeddings + outputs = self.embed_model.encode([text])[0] + return outputs def set_embedding(self, tx, entity_name, embedding): tx.run(""" @@ -328,160 +353,71 @@ class GraphDatabase: CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding) """, name=entity_name, embedding=embedding) - # def format_query_results(self, results): - # formatted_results = [] - # for row in results: - # n, rs, m = row - # entity_a = n['name'] - # entity_b = m['name'] - # for rel in rs: - # relationship = rel.type - # formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}") - # return formatted_results + def save_graph_info(self): + """ + 将图数据库的基本信息保存到工作目录中的JSON文件 + 保存的信息包括:数据库名称、状态、嵌入模型名称等 + """ + try: + # 获取数据库信息 + db_info = None + if self.status == "open" and self.driver: + try: + db_info = self.get_database_info(self.kgdb_name) + except Exception as e: + logger.warning(f"无法获取数据库信息:{e}") + # 构建要保存的信息字典 + graph_info = { + "kgdb_name": self.kgdb_name, + "status": self.status, + "embed_model_name": self.embed_model_name, + "last_updated": None, # 这里可以添加时间戳 + "database_info": db_info + } + + # 添加时间戳 + from datetime import datetime + graph_info["last_updated"] = datetime.now().isoformat() + + # 保存到JSON文件 + info_file_path = os.path.join(self.work_dir, "graph_info.json") + with open(info_file_path, 'w', encoding='utf-8') as f: + json.dump(graph_info, f, ensure_ascii=False, indent=2) + + logger.info(f"图数据库信息已保存到:{info_file_path}") + return True + except Exception as e: + logger.error(f"保存图数据库信息失败:{e}") + return False + + def load_graph_info(self): + """ + 从工作目录中的JSON文件加载图数据库的基本信息 + 返回True表示加载成功,False表示加载失败 + """ + try: + info_file_path = os.path.join(self.work_dir, "graph_info.json") + if not os.path.exists(info_file_path): + logger.warning(f"图数据库信息文件不存在:{info_file_path}") + return False + + with open(info_file_path, 'r', encoding='utf-8') as f: + graph_info = json.load(f) + + # 更新对象属性 + if graph_info.get("embed_model_name"): + self.embed_model_name = graph_info["embed_model_name"] + + # 如果需要,可以加载更多信息 + # 注意:这里不更新self.kgdb_name,因为它是在初始化时设置的 + + logger.info(f"已加载图数据库信息,最后更新时间:{graph_info.get('last_updated')}") + return True + except Exception as e: + logger.error(f"加载图数据库信息失败:{e}") + return False if __name__ == "__main__": - config = None - - kgdb_name = "neo4j" - - class EmbeddingModel(FlagModel): - def __init__(self, config, **kwargs): - - model_name_or_path = "/data2024/yyyl/model/BAAI/bge-large-zh-v1.5/" - - super().__init__(model_name_or_path, use_fp16=False, **kwargs) - - - model = EmbeddingModel(config) - # 初始化知识图谱数据库 - kgdb = GraphDatabase(config, model) - # 创建新的数据库 - # kgdb.create_graph_database("db2") - - # 返回指定数据库信息 - # info = kgdb.get_database_info(kgdb_name) - # print(info) - - # triples = [ - # { - # "h": "CCC", - # "t": "EE", - # "r": "同学" - # }, - # { - # "h": "EE", - # "t": "RR", - # "r": "同事" - # } - # ] - - # kgdb.query_by_vector("z") - # def format_query_results(results): - # formatted_results = {"nodes": [], "edges": []} - - # # 用于存储所有唯一的节点信息 - # node_dict = {} - - # for item in results: - # # 确保item[1]是一个非空的列表 - # if isinstance(item[1], list) and len(item[1]) > 0: - # relationship = item[1][0] - # rel_id = relationship.element_id - # nodes = relationship.nodes - # if len(nodes) == 2: - # node1, node2 = nodes - - # # 提取源节点和目标节点信息 - # node1_id = node1.element_id - # node2_id = node2.element_id - # node1_name = item[0] # 假设节点名称和列表中的第一个元素相同 - # node2_name = item[2] if len(item) > 2 else 'unknown' - - # # 记录节点信息 - # if node1_id not in node_dict: - # node_dict[node1_id] = {"id": node1_id, "name": node1_name} - # if node2_id not in node_dict: - # node_dict[node2_id] = {"id": node2_id, "name": node2_name} - - # # 确定关系类型 - - # relationship_type = relationship._properties.get('type', 'unknown') - # if relationship_type == 'unknown': - # relationship_type = relationship.type - - # # 记录边的信息 - # formatted_results["edges"].append({ - # "id": rel_id, - # "type": relationship_type, - # "source_id": node1_id, - # "target_id": node2_id, - # "source_name": node1_name, - # "target_name": node2_name - # }) - - # # 将唯一的节点信息添加到结果中 - # formatted_results["nodes"] = list(node_dict.values()) - - # return formatted_results - - # entities = ['jqy', '维c'] - # results = [] - # for entitie in entities: - # result = kgdb.query_by_vector(entitie) - # if result != []: - # results.extend(result) - # print(format_query_results(results)) - - - # kgdb.txt_add_vector_entity(triples, model) - # print("Extend the Graph data base") - - # kgdb.jsonl_file_add_entity("/data2024/yyyl/ProjectAthena/tep.jsonl", kgdb_name) - # print("Extend the Graph data base") - - # triples_path = "output.jsonl" - # def read_triples(file_path): - # with open(file_path, 'r', encoding='utf-8') as file: - # for line in file: - # item = json.loads(line.strip()) - # yield [item] - # for trio in read_triples(triples_path): - # kgdb.txt_add_entity(trio) - - # 通过文件添加三元组数据 - # kgdb.file_add_entity("/data2024/yyyl/ProjectAthena/test.pdf", "/data2024/yyyl/ProjectAthena/output.jsonl", kgdb_name) - # print("Extend the Graph data base") - - # 删除数据库信息 - # kgdb.delete_entity() - # print("Clear the Graph data base") - - # 查询所有节点和关系 - # results = kgdb.query_all_nodes_and_relationships(kgdb_name) - # print(results) - - # 查询特定实体及其关系 - # results = kgdb.query_specific_entity("kllll", kgdb_name) - # print(results) - - # 查询特定实体及其关系 - # results = kgdb.query_by_vector_tep("z", model, tokenizer) - # print(results) - - # 查询特定关系类型的所有节点 - # results = kgdb.query_by_relationship_type("作用或食用效果", kgdb_name) - # print(results) - - # 模糊查询 - # results = kgdb.query_entity_like("三七", kgdb_name) - # print(results) - - # 查询节点信息 - # results = kgdb.query_entity_like("三七提取物", kgdb_name) - # print(results) - - # 关闭数据库连接 - kgdb.close() - + pass \ No newline at end of file diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 99ecb830..4b16b58d 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -1,8 +1,7 @@ import os from pymilvus import MilvusClient, MilvusException -from src.utils import setup_logger, hashstr -logger = setup_logger("KnowledgeBase") +from src.utils import logger, hashstr class KnowledgeBase: @@ -68,7 +67,7 @@ class KnowledgeBase: logger.error(f"Collection {collection_name} not found, create it") # self.add_collection(collection_name) - vectors = self.embed_model.encode(docs) + vectors = self.embed_model.batch_encode(docs) data = [{ "id": int(random.random() * 1e12), diff --git a/src/core/retriever.py b/src/core/retriever.py index 1cfb862a..55d21a9b 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,7 +1,5 @@ from src.models.rerank_model import get_reranker -from src.utils.logging_config import setup_logger - -logger = setup_logger("server-common") +from src.utils.logging_config import logger class Retriever: @@ -19,7 +17,6 @@ class Retriever: self.web_searcher = WebSearcher() def retrieval(self, query, history, meta): - refs = {"query": query, "history": history, "meta": meta} refs["model_name"] = self.config.model_name refs["entities"] = self.reco_entities(query, history, refs) @@ -66,8 +63,8 @@ class Retriever: def query_classification(self, query): """判断是否需要查询 - - 对于完全基于用户给定信息的任务,称之为“足够”“sufficient”,不需要检索; - - 否则,称之为“不足”“insufficient”,可能需要检索, + - 对于完全基于用户给定信息的任务,称之为"足够""sufficient",不需要检索; + - 否则,称之为"不足""insufficient",可能需要检索, """ raise NotImplementedError @@ -107,12 +104,14 @@ class Retriever: distance_threshold = meta.get("distanceThreshold", 0) top_k = 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"]) kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] + # 重排序 if self.config.enable_reranker and len(kb_res) > 0: texts = [r["entity"]["text"] for r in kb_res] rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True) @@ -176,7 +175,7 @@ class Retriever: return entities - def _extract_relationship_info(self, relationship, source_name, target_name): + def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None): """ 提取关系信息并返回格式化的节点和边信息 """ @@ -189,6 +188,9 @@ class Retriever: source_id = source.element_id target_id = target.element_id + source_name = node_dict[source_id]["name"] if source_name is None else source_name + target_name = node_dict[target_id]["name"] if target_name is None else target_name + relationship_type = relationship._properties.get("type", "unknown") if relationship_type == "unknown": relationship_type = relationship.type @@ -230,24 +232,33 @@ class Retriever: return formatted_results def format_query_results(self, results): + logger.debug(f"Graph Query Results: {results}") formatted_results = {"nodes": [], "edges": []} node_dict = {} for item in results: - if not isinstance(item[1], list) or not item[1]: + # 检查数据格式 + if len(item) < 2 or not isinstance(item[1], list): continue - relationship = item[1][0] - source_name = item[0] - target_name = item[2] if len(item) > 2 else "unknown" + node_dict[item[0].element_id] = dict(id=item[0].element_id, name=item[0]._properties.get("name", "Unknown")) + node_dict[item[2].element_id] = dict(id=item[2].element_id, name=item[2]._properties.get("name", "Unknown")) - node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name) - if node_info is None or edge_info is None: - continue + # 处理关系列表中的每个关系 + for i, relationship in enumerate(item[1]): + try: + # 提取关系信息 + node_info, edge_info = self._extract_relationship_info(relationship, node_dict=node_dict) + if node_info is None or edge_info is None: + continue - node_dict.update({node["id"]: node for node in node_info}) - formatted_results["edges"].append(edge_info) + # 添加边 + formatted_results["edges"].append(edge_info) + except Exception as e: + logger.error(f"处理关系时出错: {e}, 关系: {relationship}") + continue + # 将节点字典转换为列表 formatted_results["nodes"] = list(node_dict.values()) return formatted_results diff --git a/src/core/startup.py b/src/core/startup.py index edab9216..11092a2d 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -3,9 +3,7 @@ from src.core import DataBaseManager from src.core.retriever import Retriever from src.models import select_model from src.config import Config -from src.utils import setup_logger - -logger = setup_logger("Startup") +from src.utils import logger class Startup: diff --git a/src/main.py b/src/main.py index d831e3d0..c5bea229 100644 --- a/src/main.py +++ b/src/main.py @@ -1,11 +1,13 @@ import uvicorn from dotenv import load_dotenv + +load_dotenv("src/.env") + from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from src.routers import router -from src.utils.logging_config import setup_logger +from src.utils.logging_config import logger -load_dotenv() app = FastAPI() app.include_router(router) @@ -20,9 +22,6 @@ app.add_middleware( ) -logger = setup_logger("server:main") - - if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=5000, threads=10, workers=10) diff --git a/src/models/__init__.py b/src/models/__init__.py index c3d0a280..e4ec33e4 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -18,6 +18,7 @@ def select_model(config, model_provider=None, model_name=None): "siliconflow", "zhipu", "lingyiwanwu", + "together.ai", ]: return OpenAIBase( api_key=os.getenv(model_info["env"][0]), diff --git a/src/models/chat_model.py b/src/models/chat_model.py index cfbcf4bb..49f1df9b 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -1,9 +1,6 @@ import os from openai import OpenAI -from src.utils.logging_config import setup_logger - - -logger = setup_logger(__name__) +from src.utils import logger class OpenAIBase(): def __init__(self, api_key, base_url, model_name): diff --git a/src/models/embedding.py b/src/models/embedding.py index 2bcd8906..853a9020 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -4,12 +4,9 @@ import requests from FlagEmbedding import FlagModel from src.config import EMBED_MODEL_INFO -from src.utils.logging_config import setup_logger -from src.utils import hashstr +from src.utils import hashstr, logger -logger = setup_logger("EmbeddingModel") - class LocalEmbeddingModel(FlagModel): def __init__(self, config, **kwargs): info = EMBED_MODEL_INFO[config.embed_model] @@ -31,6 +28,7 @@ class RemoteEmbeddingModel: embed_state = {} def batch_encode(self, messages, batch_size=20): + logger.info(f"Batch encoding {len(messages)} messages") data = [] if len(messages) > batch_size: diff --git a/src/models/ollama_embedding.py b/src/models/ollama_embedding.py index 0819a46b..1aefc16c 100644 --- a/src/models/ollama_embedding.py +++ b/src/models/ollama_embedding.py @@ -4,9 +4,7 @@ import numpy as np from typing import List, Union, Dict from src.models.embedding import RemoteEmbeddingModel -from src.utils.logging_config import setup_logger - -logger = setup_logger("OllamaEmbedding") +from src.utils.logging_config import logger class OllamaEmbedding(RemoteEmbeddingModel): """ diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index dbd24b0d..5f02a75e 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -5,10 +5,7 @@ import numpy as np from FlagEmbedding import FlagReranker from src.config import RERANKER_LIST -from src.utils.logging_config import setup_logger - - -logger = setup_logger("RerankModel") +from src.utils.logging_config import logger class LocalReranker(FlagReranker): diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index 3ca3d60b..03bceacf 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -11,8 +11,7 @@ from transformers import ( BitsAndBytesConfig ) -from src.utils import setup_logger -logger = setup_logger("OneKE") +from src.utils import logger dotenv.load_dotenv() diff --git a/src/routers/base_router.py b/src/routers/base_router.py index 7d9b0f21..0dcbbde7 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -5,11 +5,11 @@ base = APIRouter() from fastapi import FastAPI, HTTPException from fastapi.responses import JSONResponse from fastapi import Request, Body -from src.core import HistoryManager -from src.utils.logging_config import setup_logger -from src.core.startup import startup -logger = setup_logger("server-base") +from src.core import HistoryManager +from src.core.startup import startup +from src.utils import logger + @base.get("/") async def route_index(): @@ -39,6 +39,6 @@ def get_log(): last_lines = deque(f, maxlen=1000) log = ''.join(last_lines) - return {"log": log} + return {"log": log, "message": "success", "log_file": LOG_FILE} diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index f199c791..07dda597 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -5,14 +5,12 @@ from fastapi.responses import StreamingResponse, Response from concurrent.futures import ThreadPoolExecutor from src.core import HistoryManager from src.core.startup import startup -from src.utils.logging_config import setup_logger +from src.utils.logging_config import logger chat = APIRouter(prefix="/chat") -logger = setup_logger("server-chat") # 创建线程池 executor = ThreadPoolExecutor() -refs_pool = {} @chat.get("/") async def chat_get(): @@ -26,6 +24,7 @@ def chat_post( cur_res_id: str = Body(...)): history_manager = HistoryManager(history) + logger.debug(f"Received query: {query} with meta: {meta}") def make_chunk(content=None, **kwargs): return json.dumps({ @@ -112,10 +111,4 @@ async def call(query: str = Body(...), meta: dict = Body(None)): response = await predict_async(query) logger.debug({"query": query, "response": response.content}) - return {"response": response.content} - -@chat.get("/refs") -def get_refs(cur_res_id: str): - global refs_pool - refs = refs_pool.pop(cur_res_id, None) - return {"refs": refs} \ No newline at end of file + return {"response": response.content} \ No newline at end of file diff --git a/src/routers/data_router.py b/src/routers/data_router.py index 9d39d01e..6a2bf2ce 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -2,12 +2,11 @@ import os from typing import List, Optional from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body -from src.utils import setup_logger, hashstr +from src.utils import logger, hashstr from src.core.startup import startup data = APIRouter(prefix="/data") -logger = setup_logger("server-database") @data.get("/") def get_databases(): @@ -100,7 +99,6 @@ async def get_graph_info(): @data.get("/graph/node") async def get_graph_node(entity_name: str): - logger.debug(f"Get graph node {entity_name}") result = startup.dbm.graph_base.query_node(entity_name=entity_name) return {"result": startup.retriever.format_query_results(result), "message": "success"} @@ -113,7 +111,7 @@ async def get_graph_nodes(kgdb_name: str, num: int): result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num) return {"result": startup.retriever.format_general_results(result), "message": "success"} -@data.post("/graph/add") +@data.post("/graph/add-by-jsonl") async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str] = Body(None)): if not startup.config.enable_knowledge_graph: raise HTTPException(status_code=400, detail="Knowledge graph is not enabled") diff --git a/src/routers/tool_router.py b/src/routers/tool_router.py index 6732fed9..5690c086 100644 --- a/src/routers/tool_router.py +++ b/src/routers/tool_router.py @@ -3,11 +3,10 @@ from fastapi import APIRouter, Body from pydantic import BaseModel from typing import List, Dict, Any, Optional -from src.utils import setup_logger +from src.utils import logger tool = APIRouter(prefix="/tool") -logger = setup_logger("server-tools") class Tool(BaseModel): name: str diff --git a/src/static/models.yaml b/src/static/models.yaml index a2610c2c..329f9f54 100644 --- a/src/static/models.yaml +++ b/src/static/models.yaml @@ -50,6 +50,18 @@ MODEL_NAMES: - deepseek-ai/DeepSeek-R1-Distill-Qwen-7B - Qwen/Qwen2.5-72B-Instruct - Qwen/Qwen2.5-7B-Instruct + together.ai: + name: Together.ai + url: https://api.together.ai/models + base_url: https://api.together.xyz/v1/ + default: meta-llama/Llama-3.3-70B-Instruct-Turbo-Free + env: + - TOGETHER_API_KEY + models: + - meta-llama/Llama-3.3-70B-Instruct-Turbo + - meta-llama/Llama-3.3-70B-Instruct-Turbo-Free + - deepseek-ai/DeepSeek-R1-Distill-Llama-70B-free + qianfan: name: 百度千帆 (QianFan) url: https://open.bigmodel.cn/dev/api diff --git a/src/utils/__init__.py b/src/utils/__init__.py index 5b613e7a..37ff7c4a 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -1,6 +1,6 @@ import time import random -from src.utils.logging_config import setup_logger, logger +from src.utils.logging_config import logger def is_text_pdf(pdf_path): import fitz diff --git a/src/utils/web_search.py b/src/utils/web_search.py index e66ecf86..cb343ddd 100644 --- a/src/utils/web_search.py +++ b/src/utils/web_search.py @@ -1,9 +1,7 @@ import os from typing import List, Dict from tavily import TavilyClient -from src.utils.logging_config import setup_logger - -logger = setup_logger("web-search") +from src.utils.logging_config import logger class WebSearcher: def __init__(self): diff --git a/test/data/A_Dream_of_Red_Mansions_tiny.jsonl b/test/data/A_Dream_of_Red_Mansions_tiny.jsonl new file mode 100644 index 00000000..5283f7db --- /dev/null +++ b/test/data/A_Dream_of_Red_Mansions_tiny.jsonl @@ -0,0 +1,5 @@ +{"h": "贾代善", "t": "贾源", "r": "子"} +{"h": "娄氏", "t": "贾源", "r": "重孙媳妇"} +{"h": "贾母", "t": "贾代善", "r": "妻"} +{"h": "老姨奶奶", "t": "贾代善", "r": "妾"} +{"h": "贾敏", "t": "贾代善", "r": "女"} \ No newline at end of file diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 8fc760e6..70b66730 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -1,6 +1,6 @@