Skip to content

Commit 77925a9

Browse files
committed
feature:merge source-embedding
2 parents 5a5fba5 + ce4c3e8 commit 77925a9

7 files changed

Lines changed: 95 additions & 39 deletions

File tree

pilot/conversation.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -231,8 +231,8 @@ def gen_sqlgen_conversation(dbname):
231231
sep2="</s>",
232232
)
233233

234-
conv_qa_prompt_template = """ 基于以下已知的信息, 专业、详细的回答用户的问题,
235-
如果无法从提供的恶内容中获取答案, 请说: "知识库中提供的内容不足以回答此问题", 但是你可以给出一些与问题相关答案的建议
234+
conv_qa_prompt_template = """ 基于以下已知的信息, 专业、简要的回答用户的问题,
235+
如果无法从提供的恶内容中获取答案, 请说: "知识库中提供的内容不足以回答此问题" 禁止胡乱编造
236236
已知内容:
237237
{context}
238238
问题:

pilot/server/webserver.py

Lines changed: 53 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,9 @@
1111
import datetime
1212
import requests
1313
from urllib.parse import urljoin
14+
15+
from langchain import PromptTemplate
16+
1417
from pilot.configs.model_config import DB_SETTINGS, KNOWLEDGE_UPLOAD_ROOT_PATH, LLM_MODEL_CONFIG
1518
from pilot.server.vectordb_qa import KnownLedgeBaseQA
1619
from pilot.connections.mysql import MySQLOperator
@@ -32,7 +35,7 @@
3235
conv_templates,
3336
conversation_types,
3437
conversation_sql_mode,
35-
SeparatorStyle
38+
SeparatorStyle, conv_qa_prompt_template
3639
)
3740

3841
from pilot.utils import (
@@ -57,6 +60,8 @@
5760
dbs = []
5861
vs_list = ["新建知识库"] + get_vector_storelist()
5962
autogpt = False
63+
vector_store_client = None
64+
vector_store_name = {"vs_name": ""}
6065

6166
priority = {
6267
"vicuna-13b": "aaa"
@@ -106,7 +111,7 @@ def get_database_list():
106111
def load_demo(url_params, request: gr.Request):
107112
logger.info(f"load_demo. ip: {request.client.host}. params: {url_params}")
108113

109-
dbs = get_database_list()
114+
# dbs = get_database_list()
110115
dropdown_update = gr.Dropdown.update(visible=True)
111116
if dbs:
112117
gr.Dropdown.update(choices=dbs)
@@ -224,15 +229,33 @@ def http_bot(state, mode, sql_mode, db_selector, temperature, max_new_tokens, re
224229
state.messages[0][0] = ""
225230
state.messages[0][1] = ""
226231
state.messages[-2][1] = follow_up_prompt
227-
232+
prompt = state.get_prompt()
233+
skip_echo_len = len(prompt.replace("</s>", " ")) + 1
228234
if mode == conversation_types["default_knownledge"] and not db_selector:
229235
query = state.messages[-2][1]
230236
knqa = KnownLedgeBaseQA()
231237
state.messages[-2][1] = knqa.get_similar_answer(query)
232-
233-
prompt = state.get_prompt()
234-
235-
skip_echo_len = len(prompt.replace("</s>", " ")) + 1
238+
prompt = state.get_prompt()
239+
state.messages[-2][1] = query
240+
skip_echo_len = len(prompt.replace("</s>", " ")) + 1
241+
242+
if mode == conversation_types["custome"] and not db_selector:
243+
persist_dir = os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vector_store_name["vs_name"] + ".vectordb")
244+
print("向量数据库持久化地址: ", persist_dir)
245+
knowledge_embedding_client = KnowledgeEmbedding(file_path="", model_name=LLM_MODEL_CONFIG["sentence-transforms"], vector_store_config={"vector_store_name": vector_store_name["vs_name"],
246+
"vector_store_path": KNOWLEDGE_UPLOAD_ROOT_PATH})
247+
query = state.messages[-2][1]
248+
docs = knowledge_embedding_client.similar_search(query, 1)
249+
context = [d.page_content for d in docs]
250+
prompt_template = PromptTemplate(
251+
template=conv_qa_prompt_template,
252+
input_variables=["context", "question"]
253+
)
254+
result = prompt_template.format(context="\n".join(context), question=query)
255+
state.messages[-2][1] = result
256+
prompt = state.get_prompt()
257+
state.messages[-2][1] = query
258+
skip_echo_len = len(prompt.replace("</s>", " ")) + 1
236259

237260
# Make requests
238261
payload = {
@@ -438,9 +461,10 @@ def build_single_model_ui():
438461

439462
load_file_button = gr.Button("上传并加载到知识库")
440463
with gr.Tab("上传文件夹"):
441-
folder_files = gr.File(label="添加文件",
464+
folder_files = gr.File(label="添加文件夹",
465+
accept_multiple_files=True,
442466
file_count="directory",
443-
show_label=False)
467+
show_label=False)
444468
load_folder_button = gr.Button("上传并加载到知识库")
445469

446470
with gr.Blocks():
@@ -483,15 +507,17 @@ def build_single_model_ui():
483507
[state, mode, sql_mode, db_selector, temperature, max_output_tokens],
484508
[state, chatbot] + btn_list
485509
)
510+
vs_add.click(fn=save_vs_name, show_progress=True,
511+
inputs=[vs_name],
512+
outputs=[vs_name])
486513
load_file_button.click(fn=knowledge_embedding_store,
487514
show_progress=True,
488515
inputs=[vs_name, files],
489516
outputs=[vs_name])
490-
# load_folder_button.click(get_vector_store,
491-
# show_progress=True,
492-
# inputs=[vs_name, folder_files, 100 , chatbot, vs_add,
493-
# vs_add],
494-
# outputs=["db-out", folder_files, chatbot])
517+
load_folder_button.click(fn=knowledge_embedding_store,
518+
show_progress=True,
519+
inputs=[vs_name, folder_files],
520+
outputs=[vs_name])
495521
return state, chatbot, textbox, send_btn, button_row, parameter_row
496522

497523

@@ -531,17 +557,26 @@ def build_webdemo():
531557
return demo
532558

533559

560+
def save_vs_name(vs_name):
561+
vector_store_name["vs_name"] = vs_name
562+
return vs_name
563+
534564
def knowledge_embedding_store(vs_id, files):
535565
# vs_path = os.path.join(VS_ROOT_PATH, vs_id)
536566
if not os.path.exists(os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id)):
537567
os.makedirs(os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id))
538568
for file in files:
539569
filename = os.path.split(file.name)[-1]
540570
shutil.move(file.name, os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id, filename))
571+
knowledge_embedding_client = KnowledgeEmbedding(
572+
file_path=os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id, filename),
573+
model_name=LLM_MODEL_CONFIG["sentence-transforms"],
574+
vector_store_config={
575+
"vector_store_name": vector_store_name["vs_name"],
576+
"vector_store_path": KNOWLEDGE_UPLOAD_ROOT_PATH})
577+
knowledge_embedding_client.knowledge_embedding()
578+
541579

542-
knowledge_embedding = KnowledgeEmbedding.knowledge_embedding(os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id, filename), LLM_MODEL_CONFIG["sentence-transforms"], {"vector_store_name": vs_id,
543-
"vector_store_path": KNOWLEDGE_UPLOAD_ROOT_PATH})
544-
knowledge_embedding.source_embedding()
545580
logger.info("knowledge embedding success")
546581
return os.path.join(KNOWLEDGE_UPLOAD_ROOT_PATH, vs_id, vs_id + ".vectordb")
547582

@@ -558,7 +593,7 @@ def knowledge_embedding_store(vs_id, files):
558593
args = parser.parse_args()
559594
logger.info(f"args: {args}")
560595

561-
dbs = get_database_list()
596+
# dbs = get_database_list()
562597

563598
# 加载插件
564599
cfg = Config()

pilot/source_embedding/csv_embedding.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ class CSVEmbedding(SourceEmbedding):
1010

1111
def __init__(self, file_path, model_name, vector_store_config, embedding_args: Optional[Dict] = None):
1212
"""Initialize with csv path."""
13+
super().__init__(file_path, model_name, vector_store_config)
1314
self.file_path = file_path
1415
self.model_name = model_name
1516
self.vector_store_config = vector_store_config

pilot/source_embedding/knowledge_embedding.py

Lines changed: 26 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -4,17 +4,31 @@
44

55

66
class KnowledgeEmbedding:
7-
@staticmethod
8-
def knowledge_embedding(file_path:str, model_name, vector_store_config):
9-
if file_path.endswith(".pdf"):
10-
embedding = PDFEmbedding(file_path=file_path, model_name=model_name,
11-
vector_store_config=vector_store_config)
12-
elif file_path.endswith(".md"):
13-
embedding = MarkdownEmbedding(file_path=file_path, model_name=model_name,
14-
vector_store_config=vector_store_config)
7+
def __init__(self, file_path, model_name, vector_store_config):
8+
"""Initialize with Loader url, model_name, vector_store_config"""
9+
self.file_path = file_path
10+
self.model_name = model_name
11+
self.vector_store_config = vector_store_config
12+
self.vector_store_type = "default"
13+
self.knowledge_embedding_client = self.init_knowledge_embedding()
1514

16-
elif file_path.endswith(".csv"):
17-
embedding = CSVEmbedding(file_path=file_path, model_name=model_name,
18-
vector_store_config=vector_store_config)
15+
def knowledge_embedding(self):
16+
self.knowledge_embedding_client.source_embedding()
1917

20-
return embedding
18+
def init_knowledge_embedding(self):
19+
if self.file_path.endswith(".pdf"):
20+
embedding = PDFEmbedding(file_path=self.file_path, model_name=self.model_name,
21+
vector_store_config=self.vector_store_config)
22+
elif self.file_path.endswith(".md"):
23+
embedding = MarkdownEmbedding(file_path=self.file_path, model_name=self.model_name, vector_store_config=self.vector_store_config)
24+
25+
elif self.file_path.endswith(".csv"):
26+
embedding = CSVEmbedding(file_path=self.file_path, model_name=self.model_name,
27+
vector_store_config=self.vector_store_config)
28+
elif self.vector_store_type == "default":
29+
embedding = MarkdownEmbedding(file_path=self.file_path, model_name=self.model_name, vector_store_config=self.vector_store_config)
30+
31+
return embedding
32+
33+
def similar_search(self, text, topk):
34+
return self.knowledge_embedding_client.similar_search(text, topk)

pilot/source_embedding/markdown_embedding.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@ class MarkdownEmbedding(SourceEmbedding):
1515

1616
def __init__(self, file_path, model_name, vector_store_config):
1717
"""Initialize with markdown path."""
18+
super().__init__(file_path, model_name, vector_store_config)
1819
self.file_path = file_path
1920
self.model_name = model_name
2021
self.vector_store_config = vector_store_config

pilot/source_embedding/pdf_embedding.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,9 +13,12 @@ class PDFEmbedding(SourceEmbedding):
1313

1414
def __init__(self, file_path, model_name, vector_store_config):
1515
"""Initialize with pdf path."""
16+
super().__init__(file_path, model_name, vector_store_config)
1617
self.file_path = file_path
1718
self.model_name = model_name
1819
self.vector_store_config = vector_store_config
20+
# SourceEmbedding(file_path =file_path, );
21+
SourceEmbedding(file_path, model_name, vector_store_config)
1922

2023
@register
2124
def read(self):

pilot/source_embedding/source_embedding.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -22,12 +22,16 @@ class SourceEmbedding(ABC):
2222
Implementations should implement the method
2323
"""
2424

25-
def __init__(self, yuque_path, model_name, vector_store_config, embedding_args: Optional[Dict] = None):
26-
"""Initialize with YuqueLoader url, model_name, vector_store_config"""
27-
self.yuque_path = yuque_path
25+
def __init__(self, file_path, model_name, vector_store_config, embedding_args: Optional[Dict] = None):
26+
"""Initialize with Loader url, model_name, vector_store_config"""
27+
self.file_path = file_path
2828
self.model_name = model_name
2929
self.vector_store_config = vector_store_config
3030
self.embedding_args = embedding_args
31+
self.embeddings = HuggingFaceEmbeddings(model_name=self.model_name)
32+
persist_dir = os.path.join(self.vector_store_config["vector_store_path"],
33+
self.vector_store_config["vector_store_name"] + ".vectordb")
34+
self.vector_store_client = Chroma(persist_directory=persist_dir, embedding_function=self.embeddings)
3135

3236
@abstractmethod
3337
@register
@@ -50,18 +54,16 @@ def text_to_vector(self, docs):
5054
@register
5155
def index_to_store(self, docs):
5256
"""index to vector store"""
53-
embeddings = HuggingFaceEmbeddings(model_name=self.model_name)
54-
5557
persist_dir = os.path.join(self.vector_store_config["vector_store_path"],
5658
self.vector_store_config["vector_store_name"] + ".vectordb")
57-
self.vector_store = Chroma.from_documents(docs, embeddings, persist_directory=persist_dir)
59+
self.vector_store = Chroma.from_documents(docs, self.embeddings, persist_directory=persist_dir)
5860
self.vector_store.persist()
5961

6062
@register
6163
def similar_search(self, doc, topk):
6264
"""vector store similarity_search"""
6365

64-
return self.vector_store.similarity_search(doc, topk)
66+
return self.vector_store_client.similarity_search(doc, topk)
6567

6668
def source_embedding(self):
6769
if 'read' in registered_methods:

0 commit comments

Comments
 (0)