From 6f81b061d2b626dcd8ce88ee0a4c293d91975c5a Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 25 Sep 2024 13:46:23 +0800 Subject: [PATCH] update tools view and indexing --- src/core/database.py | 40 ++-- src/core/indexing.py | 37 ++++ src/core/retriever.py | 3 +- src/utils/prompts.py | 16 ++ src/views/__init__.py | 2 + src/views/tools_view.py | 45 ++++ web/src/assets/base.css | 43 ++-- web/src/components/TextChunkingComponent.vue | 211 +++++++++++++++++++ web/src/views/ToolsView.vue | 123 ++++++----- 9 files changed, 429 insertions(+), 91 deletions(-) create mode 100644 src/core/indexing.py create mode 100644 src/views/tools_view.py create mode 100644 web/src/components/TextChunkingComponent.vue diff --git a/src/core/database.py b/src/core/database.py index 440e2300..3c5cd5d4 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -2,7 +2,6 @@ import os import json import time from src.utils import hashstr, setup_logger, is_text_pdf -from src.core.filereader import pdfreader, plainreader from src.models.embedding import get_embedding_model logger = setup_logger("DataBaseManager") @@ -95,18 +94,16 @@ class DataBaseManager: self._save_databases() return self.get_databases() - def add_files(self, db_id, files): + def add_files(self, db_id, files, params=None): db = self.get_kb_by_id(db_id) if db.embed_model != self.config.embed_model: logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}") return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"} + # Preprocessing the files to the queue new_files = [] for file in files: - # filenames = [f["filename"] for f in db.files] - # if os.path.basename(file) in filenames: - # continue new_file = { "file_id": "file_" + hashstr(file + str(time.time())), "filename": os.path.basename(file), @@ -118,25 +115,30 @@ class DataBaseManager: db.files.append(new_file) new_files.append(new_file) + from src.core.indexing import chunk_file, chunk_text for new_file in new_files: file_id = new_file["file_id"] - idx = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0] + idx = self.get_idx_by_fileid(db, file_id) db.files[idx]["status"] = "processing" try: - text = self.read_text(new_file["path"]) - chunks = self.chunking(text) + if new_file["type"] in ["txt", "docx", "md"]: + nodes = chunk_file(new_file["path"], params=params) + else: + text = self.read_text(new_file["path"]) + nodes = chunk_text(text, params) + self.knowledge_base.add_documents( - docs=chunks, + docs=[node.to_dict() for node in nodes], collection_name=db.metaname, file_id=file_id) - idx = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0] + idx = self.get_idx_by_fileid(db, file_id) db.files[idx]["status"] = "done" except Exception as e: logger.error(f"Failed to add documents to collection {db.metaname}, {e}") - idx = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0] + idx = self.get_idx_by_fileid(db, file_id) db.files[idx]["status"] = "failed" self._save_databases() @@ -151,7 +153,7 @@ class DataBaseManager: db.update(self.knowledge_base.get_collection_info(db.metaname)) return db.to_dict() - def read_text(self, file): + def read_text(self, file, params=None): support_format = [".pdf", ".txt", ".md"] assert os.path.exists(file), "File not found" logger.info(f"Try to read file {file}") @@ -162,12 +164,14 @@ class DataBaseManager: if file.endswith(".pdf"): if is_text_pdf(file): + from src.core.filereader import pdfreader return pdfreader(file) else: from src.plugins import pdf2txt return pdf2txt(file, return_text=True) elif file.endswith(".txt") or file.endswith(".md"): + from src.core.filereader import plainreader return plainreader(file) else: @@ -176,7 +180,7 @@ class DataBaseManager: def delete_file(self, db_id, file_id): db = self.get_kb_by_id(db_id) - file_idx_to_delete = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0] + file_idx_to_delete = self.get_idx_by_fileid(db, file_id) self.knowledge_base.client.delete( collection_name=db.metaname, @@ -196,7 +200,10 @@ class DataBaseManager: ) return {"lines": lines} - def chunking(self, text, chunk_size=1024): + def chunking(self, text, params=None): + chunk_method = params.get("chunk_method", "fixed") + chunk_size = params.get("chunk_size", 500) + """将文本切分成固定大小的块""" chunks = [] for i in range(0, len(text), chunk_size): @@ -219,6 +226,11 @@ class DataBaseManager: return db return None + def get_idx_by_fileid(self, db, file_id): + for idx, f in enumerate(db.files): + if f["file_id"] == file_id: + return idx + class DataBaseLite: def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None: diff --git a/src/core/indexing.py b/src/core/indexing.py new file mode 100644 index 00000000..eb264a8a --- /dev/null +++ b/src/core/indexing.py @@ -0,0 +1,37 @@ +import os +from pathlib import Path +from llama_index.core.node_parser import SimpleFileNodeParser +from llama_index.readers.file import FlatReader + +from src.utils import hashstr + +def chunk_text(text, params=None): + params = params or {} + from llama_index.core import Document + from llama_index.core.node_parser import SentenceSplitter + chunk_size = int(params.get("chunk_size", 500)) + chunk_overlap = int(params.get("chunk_overlap", 20)) + splitter = SentenceSplitter( + chunk_size=chunk_size, + chunk_overlap=chunk_overlap, + ) + doc = Document(id_=hashstr(text), text=text) + nodes = splitter.get_nodes_from_documents([doc]) + return nodes + +def chunk_file(file, params=None): + parser = SimpleFileNodeParser() + if file.endswith(".txt"): + from llama_index.readers.file import FlatReader + docs = FlatReader().load_data(Path(file)) + elif file.endswith(".docx"): + from llama_index.readers.file import DocxReader + docs = DocxReader().load_data(Path(file)) + elif file.endswith(".md"): + from llama_index.readers.file import MarkdownReader + docs = MarkdownReader().load_data(Path(file)) + else: + raise ValueError("Unsupported file type") + + nodes = parser.get_nodes_from_documents(docs) + return nodes # 返回节点 \ No newline at end of file diff --git a/src/core/retriever.py b/src/core/retriever.py index 788b4107..3d21971e 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -45,9 +45,10 @@ class Retriever: external_parts.extend(["图数据库信息:", db_text]) # 构造查询 + from src.utils.prompts import knowbase_qa_template if external_parts and len(external_parts) > 0: external = "\n\n".join(external_parts) - query = f"参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答问题。\n\n问题:{query}\n\n回答:" + query = knowbase_qa_template.format(external=external, query=query) return query diff --git a/src/utils/prompts.py b/src/utils/prompts.py index 208d067d..f105ce4c 100644 --- a/src/utils/prompts.py +++ b/src/utils/prompts.py @@ -1,3 +1,19 @@ +system_prompt = """ +""" + + +knowbase_qa_template = """ +请利用查询到的资料回答问题,回答问题时,不要过度的分点作答。如果非要分点作答,可以使用 一、二、等: + +<参考资料>: +{external} + + +<问题> +{query} +" +""" + rewritten_query_prompt_template = """ <指令>根据提供的历史信息对问题进行优化和改写,返回的问题必须符合以下内容要求和格式要求。严格不能出现禁止内容<指令> <禁止>1.绝对不能自己编造无关内容,若不能改写或无需改写直接返回原本问题 diff --git a/src/views/__init__.py b/src/views/__init__.py index 6886b6aa..5fa4b936 100644 --- a/src/views/__init__.py +++ b/src/views/__init__.py @@ -2,6 +2,7 @@ from flask import Flask from flask_cors import CORS from src.views.common_view import common from src.views.database_view import db +from src.views.tools_view import tools def create_app(): @@ -10,5 +11,6 @@ def create_app(): app.register_blueprint(common) app.register_blueprint(db) + app.register_blueprint(tools) return app diff --git a/src/views/tools_view.py b/src/views/tools_view.py new file mode 100644 index 00000000..c0466d1b --- /dev/null +++ b/src/views/tools_view.py @@ -0,0 +1,45 @@ +import os +import json +import threading +from flask import Blueprint, jsonify, request, Response + +from src.utils import setup_logger, hashstr +from src.core.startup import startup + +tools = Blueprint('tools', __name__, url_prefix="/tools") + +logger = setup_logger("server-tools") + + +@tools.route("/", methods=["GET"]) +def route_index(): + tools = [ + { + "name": "text_chunking", + "title": "Text Chunking", + "description": "Chunking text into smaller pieces for better understanding.", + "url": "/tools/text_chunking", + "method": "POST", + "params": [ + { + "name": "text", + "type": "string", + "description": "Text to be chunked." + }, + { + "name": "chunk_size", + "type": "int", + } + ] + } + ] + + return jsonify(tools) + + +@tools.route("/text_chunking", methods=["POST"]) +def text_chunking(): + from src.core.indexing import chunk_text + text = request.json.get("text") + nodes = chunk_text(text, params=request.json) + return jsonify({"nodes": [node.to_dict() for node in nodes]}) diff --git a/web/src/assets/base.css b/web/src/assets/base.css index 025a5ebf..b54140ef 100644 --- a/web/src/assets/base.css +++ b/web/src/assets/base.css @@ -1,5 +1,7 @@ /* color palette from */ +/* https://material-foundation.github.io/material-theme-builder/ */ :root { + --main-1000: #002A36; --main-900: #003A51; --main-800: #004F69; --main-700: #00637F; @@ -12,26 +14,22 @@ --main-50: #CDF5FF; --main-25: #E6FAFF; --main-10: #F5FDFF; + --main-5: #FAFCFD; - --c-white: #ffffff; - --c-white-soft: #f8f8f8; - --c-white-mute: #f2f2f2; + --gray-2000: #0C1214; + --gray-1000: #171C1F; + --gray-900: #212729; + --gray-800: #42484A; + --gray-700: #616161; + --gray-600: #8C9194; + --gray-500: #A7ACAF; + --gray-400: #C2C7CA; + --gray-300: #DEE3E6; + --gray-200: #EDF1F4; + --gray-100: #F5FAFD; + --gray-50: #F9FDFF; - --c-black: #202428; - --c-black-soft: #222222; - --c-black-mute: #282828; - - --c-black-light-1: #333333; - --c-black-light-2: #454545; - --c-black-light-3: #666666; - --c-black-light-4: #999999; - - --c-text-light-1: var(--c-black); - --c-text-dark-1: #0D0D0D; - --c-text-dark-2: #b8b8b8; - --color-text: var(--c-black); - - --main-color: var(--main-700); + --main-color: #1c6586; --main-color-dark: #004d5c; --main-light-1: #0076AB; --main-light-2: #DAEAED; @@ -39,8 +37,13 @@ --main-light-4: #F2F5F5; --main-light-5: #F7FAFB; --main-light-6: #FAFDFD; + --secondry-color: #4e616d; + --error-color: #ba1a1a; + + --bg-sider: var(--main-5); + --color-text: var(--c-black); + --min-width: 400px; - --error-color: #f50a0d; } *, @@ -55,7 +58,7 @@ body { display: flow-root; min-height: 100vh; - color: var(--color-text); + color: var(--gray-900); line-height: 1.6; font-family: 'Roboto', 'Noto Sans SC', 'HarmonyOS Sans SC', -apple-system, BlinkMacSystemFont, 'Segoe UI', 'Helvetica Neue', Arial, sans-serif; font-size: 15px; diff --git a/web/src/components/TextChunkingComponent.vue b/web/src/components/TextChunkingComponent.vue new file mode 100644 index 00000000..2386a4b4 --- /dev/null +++ b/web/src/components/TextChunkingComponent.vue @@ -0,0 +1,211 @@ + + + + + \ No newline at end of file diff --git a/web/src/views/ToolsView.vue b/web/src/views/ToolsView.vue index 207f3caa..bdc33fa3 100644 --- a/web/src/views/ToolsView.vue +++ b/web/src/views/ToolsView.vue @@ -7,12 +7,14 @@
-
-
- +
+
+
+ +
+

{{ tool.title }}

-

{{ tool.name }}

{{ tool.description }}

@@ -21,39 +23,42 @@