1111import datetime
1212import requests
1313from urllib .parse import urljoin
14+
15+ from langchain import PromptTemplate
16+
1417from pilot .configs .model_config import DB_SETTINGS , KNOWLEDGE_UPLOAD_ROOT_PATH , LLM_MODEL_CONFIG
1518from pilot .server .vectordb_qa import KnownLedgeBaseQA
1619from pilot .connections .mysql import MySQLOperator
3235 conv_templates ,
3336 conversation_types ,
3437 conversation_sql_mode ,
35- SeparatorStyle
38+ SeparatorStyle , conv_qa_prompt_template
3639)
3740
3841from pilot .utils import (
5760dbs = []
5861vs_list = ["新建知识库" ] + get_vector_storelist ()
5962autogpt = False
63+ vector_store_client = None
64+ vector_store_name = {"vs_name" : "" }
6065
6166priority = {
6267 "vicuna-13b" : "aaa"
@@ -106,7 +111,7 @@ def get_database_list():
106111def 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+
534564def 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 ()
0 commit comments