This commit is contained in:
Wenjie Zhang 2024-08-11 17:48:39 +08:00
commit 14d2d6594f
16 changed files with 135 additions and 76 deletions

View File

@ -49,7 +49,7 @@ class Config(SimpleConfig):
# 模型配置
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
self.add_item("model_provider", default="qianfan", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"])
self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"])
self.add_item("model_name", default=None, des="模型名称")
self.add_item("embed_model", default="bge-large-zh-v1.5", des="Embedding 模型", choices=["bge-large-zh-v1.5", "zhipu"])
self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"])

View File

@ -2,10 +2,7 @@ import os
import json
import time
from src.utils import hashstr, setup_logger, is_text_pdf
from src.plugins import pdf2txt
from src.core.knowledgebase import KnowledgeBase
from src.core.filereader import pdfreader, plainreader
from src.core.graphbase import GraphDatabase
from src.models.embedding import get_embedding_model
logger = setup_logger("DataBaseManager")
@ -57,6 +54,8 @@ class DataBaseManager:
self.embed_model = get_embedding_model(config)
if self.config.enable_knowledge_base:
from src.core.knowledgebase import KnowledgeBase
from src.core.graphbase import GraphDatabase
self.knowledge_base = KnowledgeBase(config, self.embed_model)
self.graph_base = GraphDatabase(self.config, self.embed_model)
self.graph_base.start()
@ -180,6 +179,7 @@ class DataBaseManager:
if is_text_pdf(file):
return pdfreader(file)
else:
from src.plugins import pdf2txt
return pdf2txt(file, return_text=True)
elif file.endswith(".txt") or file.endswith(".md"):

View File

@ -1,14 +1,13 @@
import os
from pathlib import Path
from llama_index.readers.file import PDFReader
def pdfreader(file_path):
"""读取PDF文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
assert file_path.endswith(".pdf"), "File format not supported"
from llama_index.readers.file import PDFReader
doc = PDFReader().load_data(file=Path(file_path))
# 简单的拼接起来之后返回纯文本

View File

@ -25,9 +25,13 @@ class HistoryManager():
self.add_ai(content)
return self.messages
def get_history_with_msg(self, msg, role="user"):
def get_history_with_msg(self, msg, role="user", max_rounds=None):
"""Get history with new message, but not append it to history."""
history = self.messages[:]
if max_rounds is None:
history = self.messages[:]
else:
history = self.messages[-(2*max_rounds):]
history.append({"role": role, "content": msg})
return history

View File

@ -14,6 +14,7 @@ class KnowledgeBase:
assert embed_model, "embed_model=None"
self.embed_model = embed_model
self.client = MilvusClient(self.milvus_path)
def _init_config(self, config):
@ -72,15 +73,17 @@ class KnowledgeBase:
def search(self, query, collection_name, limit=3):
query_vectors = self.embed_model.encode_queries([query])
return self.search_by_vector(query_vectors[0], collection_name, limit)
def search_by_vector(self, vector, collection_name, limit=3):
res = self.client.search(
collection_name=collection_name, # target collection
data=query_vectors, # query vectors
data=[vector], # query vectors
limit=limit, # number of returned entities
output_fields=["text", "file_id"], # specifies fields to be returned
)
return res[0] # 因为 query 只有一个
return res[0]
def examples(self, collection_name, limit=20):
res = self.client.query(

View File

@ -15,9 +15,7 @@ class Retriever:
def retrieval(self, query, history, meta):
refs = {}
refs["meta"] = meta
refs["rewritten_query"] = self.rewrite_query(query, history, refs)
refs = {"query": query, "history": history, "meta": meta}
refs["entities"] = self.reco_entities(query, history, refs)
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
refs["graph_base"] = self.query_graph(query, history, refs)
@ -68,37 +66,47 @@ class Retriever:
def query_knowledgebase(self, query, history, refs):
"""查询知识库"""
query = refs.get("rewritten_query", query)
rw_query = self.rewrite_query(query, history, refs)
kb_res = []
final_res = []
if refs["meta"].get("db_name") and self.config.enable_knowledge_base:
db_name = refs["meta"]["db_name"]
kb = self.dbm.metaname2db[refs["meta"]["db_name"]]
kb = self.dbm.metaname2db[db_name]
limit = refs["meta"].get("queryCount", 10)
kb_res = self.dbm.knowledge_base.search(query, db_name, limit=limit)
kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit)
for r in kb_res:
r["file"] = kb.id2file(r["entity"]["file_id"])
if self.config.enable_reranker:
RERANK_THRESHOLD = 0.1
for r in kb_res:
r["rerank_score"] = self.reranker.compute_score([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"] > 0.1]
final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD]
else:
final_res = kb_res[:5]
return {"results": final_res, "all_results": kb_res}
return {"results": final_res, "all_results": kb_res, "rw_query": rw_query}
def rewrite_query(self, query, history, refs):
"""重写查询"""
if refs["meta"].get("rewrite_query") is None or history == []:
rewrite_query_span = refs["meta"].get("rewrite_query", None)
if rewrite_query_span is None or rewrite_query_span == "OFF":
rewritten_query = query
else:
from src.utils.prompts import rewritten_query_prompt_template
rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query)
history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else ""
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":
hy_doc = self.model.predict(rewritten_query).content
rewritten_query = f"{rewritten_query} {hy_doc}"
return rewritten_query
def reco_entities(self, query, history, refs):

View File

@ -2,12 +2,12 @@
### 1. 对话模型支持
模型仅支持通过API调用的模型如果是需要运行本地模型则建议使用 vllm 转成 API 服务之后使用。
模型仅支持通过API调用的模型如果是需要运行本地模型则建议使用 vllm 转成 API 服务之后使用。使用前请配置 APIKEY 后使用,配置项目参考:[.env.template](../.env.template)
|模型供应商(`config.model_provider`)|默认模型(`config.model_name`)|配置项目(`.env`)|
|:-|:-|:-|
|`qianfan`|`ernie_speed`|`QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY`|
|`zhipu`|`glm-4`|`ZHIPUAPI`|
|`zhipu`(default)|`glm-4`|`ZHIPUAPI`|
|`deepseek`|`deepseek-chat`|`DEEPSEEKAPI`|
|`vllm`|`vllm`|`VLLM_API_KEY`, `VLLM_API_BASE`|

View File

@ -64,7 +64,6 @@ class VLLM(OpenAIBase):
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
import qianfan
class GeneralResponse:
@ -76,6 +75,7 @@ class GeneralResponse:
class Qianfan:
def __init__(self, model_name="ernie_speed") -> None:
import qianfan
self.model_name = model_name
access_key = os.getenv("QIANFAN_ACCESS_KEY")
secret_key = os.getenv("QIANFAN_SECRET_KEY")

View File

@ -1,38 +1,19 @@
import os
import fitz # fitz就是pip install PyMuPDF
import cv2
from paddleocr import PPStructure, save_structure_res
from paddleocr.ppstructure.recovery.recovery_to_doc import sorted_layout_boxes, convert_info_docx
from copy import deepcopy
from tqdm import tqdm
def pdf2txt(pdf_path, return_text=False):
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])
os.makedirs(output_dir, exist_ok=True)
table_engine = PPStructure(recovery=True, lang='ch')
imgs = []
img_dir = os.path.join(output_dir, 'imgs')
if not os.path.exists(img_dir):
os.makedirs(img_dir)
pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100):
page = pdfDoc[pg]
rotate = int(0)
zoom_x = 2
zoom_y = 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_filename = os.path.join(img_dir, f'images_{pg+1}.png')
pix.save(img_filename) # os.sep
imgs.append(img_filename)
else:
img_names = sorted(os.listdir(img_dir))
imgs = [os.path.join(img_dir, img_name) for img_name in img_names]
imgs = convert_imgs(pdf_path, output_dir)
respath = os.path.join(output_dir, 'res.txt')
text = []
@ -59,7 +40,7 @@ def pdf2txt(pdf_path, return_text=False):
continue # 如果不是字典或者缺少 'text' 键,跳过当前循环
text.append('\n')
whole_text = '\n'.join(text)
whole_text = ''.join(text)
with open(respath, 'w', encoding='utf-8') as f:
f.write(whole_text)
@ -68,6 +49,29 @@ def pdf2txt(pdf_path, return_text=False):
return respath
def convert_imgs(pdf_path, output_dir):
imgs = []
img_dir = os.path.join(output_dir, 'imgs')
if not os.path.exists(img_dir):
os.makedirs(img_dir)
pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100):
page = pdfDoc[pg]
rotate = int(0)
zoom_x = 2
zoom_y = 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_filename = os.path.join(img_dir, f'images_{pg+1}.png')
pix.save(img_filename) # os.sep
imgs.append(img_filename)
else:
img_names = sorted(os.listdir(img_dir))
imgs = [os.path.join(img_dir, img_name) for img_name in img_names]
return imgs
if __name__ == "__main__":
pdf_path = r'data/file/焙烤食品工艺学.pdf'
pdf_path = r'saves/data/uploads/2e04d5_保健食品.pdf'
print(pdf2txt(pdf_path))

View File

@ -1,9 +1,6 @@
FlagEmbedding==1.2.10
Flask==3.0.3
Flask_Cors==4.0.1
llama_index==0.10.53
openai==1.35.10
pymilvus==2.4.4
python-dotenv==1.0.1
PyYAML==6.0.1
qianfan==0.4.0.1

View File

@ -34,7 +34,7 @@ def chat():
new_query, refs = startup.retriever(query, history_manager.messages, meta)
messages = history_manager.get_history_with_msg(new_query)
messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round'))
history_manager.add_user(query)
logger.debug(f"Web history: {history_manager}")

View File

@ -45,7 +45,6 @@
--main-light-6: #FAFDFD;
--min-width: 400px;
--min-header-width: 80px;
--min-sider-width: 100px;
--error-color: #f50a0d;
}

View File

@ -39,11 +39,11 @@
</a-menu>
</template>
</a-dropdown>
<div class="nav-btn text" @click="opts.showPanel = !opts.showPanel">
<div class="nav-btn text " @click="opts.showPanel = !opts.showPanel">
<component :is="opts.showPanel ? FolderOpenOutlined : FolderOutlined" /> <span class="text">选项</span>
</div>
<div v-if="opts.showPanel" class="my-panal" ref="panel">
<div class="graphbase flex-center" v-if="configStore.config.enable_knowledge_base">
<div v-if="opts.showPanel" class="my-panal swing-in-top-fwd" ref="panel">
<div class="flex-center" v-if="configStore.config.enable_knowledge_base">
知识库
<div @click.stop>
<a-dropdown>
@ -65,21 +65,24 @@
</a-dropdown>
</div>
</div>
<div class="graphbase flex-center" @click="meta.use_graph = !meta.use_graph" v-if="configStore.config.enable_knowledge_base">
<div class="flex-center" @click="meta.use_graph = !meta.use_graph" v-if="configStore.config.enable_knowledge_base">
图数据库 <div @click.stop><a-switch v-model:checked="meta.use_graph" /></div>
</div>
<div class="graphbase flex-center" @click="meta.use_web = !meta.use_web" v-if="configStore.config.enable_search_engine">
<div class="flex-center" @click="meta.use_web = !meta.use_web" v-if="configStore.config.enable_search_engine">
搜索引擎Bing <div @click.stop><a-switch v-model:checked="meta.use_web" /></div>
</div>
<div class="graphbase flex-center" @click="meta.rewrite_query = !meta.rewrite_query" v-if="configStore.config.enable_reranker">
<div class="flex-center" @click="meta.rewrite_query = !meta.rewrite_query" v-if="configStore.config.enable_reranker">
重写查询 <div @click.stop><a-switch v-model:checked="meta.rewrite_query" /></div>
</div>
<div class="graphbase flex-center" @click="meta.rewrite_query = !meta.rewrite_query">
<div class="flex-center" @click="meta.rewrite_query = !meta.rewrite_query">
流式输出 <div @click.stop><a-switch v-model:checked="meta.stream" /></div>
</div>
<div class="graphbase flex-center" @click="meta.summary_title = !meta.summary_title">
<div class="flex-center" @click="meta.summary_title = !meta.summary_title">
总结对话标题 <div @click.stop><a-switch v-model:checked="meta.summary_title" /></div>
</div>
<div class="flex-center">
最大历史轮数 <a-input-number id="inputNumber" v-model:value="meta.history_round" :min="1" :max="50" />
</div>
</div>
</div>
</div>
@ -134,8 +137,18 @@
</div>
<div v-for="(res, idx) in results" :key="idx" class="result-item">
<p class="result-id"><strong>ID:</strong> #{{ res.id }}</p>
<p class="result-distance"><strong>相似度距离:</strong> {{ res.distance }}</p>
<p class="result-rerank-score"><strong>重排序分数:</strong> {{ res.rerank_score }}</p>
<p class="result-distance">
<strong>相似度距离:</strong>
<div class=scorebar>
<a-progress :percent="(res.distance * 100).toFixed(2)" stroke-color="#1677FF" :size="[200, 10]"/>
</div>
</p>
<p class="result-rerank-score">
<strong>重排序分数:</strong>
<div class=scorebar>
<a-progress :percent="(res.rerank_score * 100).toFixed(2)" stroke-color="#1677FF" :size="[200, 10]"/>
</div>
</p>
<p class="result-text">{{ res.entity.text }}</p>
</div>
</a-drawer>
@ -216,6 +229,7 @@ const meta = reactive(JSON.parse(localStorage.getItem('meta')) || {
selectedKB: null,
stream: true,
summary_title: true,
history_round: 5,
})
// marked https://marked.js.org/
@ -517,9 +531,6 @@ watch(
justify-content: space-between;
align-items: center;
gap: 10px;
}
.graphbase {
padding: 8px 16px;
border-radius: 8px;
cursor: pointer;
@ -663,6 +674,18 @@ watch(
.result-text {
margin: 5px 0;
}
.scorebar {
margin-left: 10px;
display: inline-block;
width: 200px;
padding-bottom: 2px;
& > * {
margin: 0;
}
}
}
}
@ -780,6 +803,10 @@ button:disabled {
border-radius: 4px;
}
.slide-out-left{-webkit-animation:slide-out-left .5s cubic-bezier(.55,.085,.68,.53) both;animation:slide-out-left .5s cubic-bezier(.55,.085,.68,.53) both}
.swing-in-top-fwd {-webkit-animation: swing-in-top-fwd 0.5s cubic-bezier(0.175, 0.885, 0.320, 1.275) both;animation: swing-in-top-fwd 0.5s cubic-bezier(0.175, 0.885, 0.320, 1.275) both;}
@-webkit-keyframes swing-in-top-fwd{0%{-webkit-transform:rotateX(-100deg);transform:rotateX(-100deg);-webkit-transform-origin:top;transform-origin:top;opacity:0}100%{-webkit-transform:rotateX(0deg);transform:rotateX(0deg);-webkit-transform-origin:top;transform-origin:top;opacity:1}}@keyframes swing-in-top-fwd{0%{-webkit-transform:rotateX(-100deg);transform:rotateX(-100deg);-webkit-transform-origin:top;transform-origin:top;opacity:0}100%{-webkit-transform:rotateX(0deg);transform:rotateX(0deg);-webkit-transform-origin:top;transform-origin:top;opacity:1}}
@-webkit-keyframes slide-out-left{0%{-webkit-transform:translateX(0);transform:translateX(0);opacity:1}100%{-webkit-transform:translateX(-1000px);transform:translateX(-1000px);opacity:0}}@keyframes slide-out-left{0%{-webkit-transform:translateX(0);transform:translateX(0);opacity:1}100%{-webkit-transform:translateX(-1000px);transform:translateX(-1000px);opacity:0}}
@media (max-width: 520px) {
.chat {

View File

@ -1,6 +1,6 @@
<template>
<div class="chat-container">
<div v-if="state.isSidebarOpen" class="conversations">
<div class="conversations" :class="['conversations', { 'is-open': state.isSidebarOpen }]">
<div class="actions">
<!-- <div class="action new" @click="addNewConv"><FormOutlined /></div> -->
<span style="font-weight: bold;">对话历史</span>
@ -120,20 +120,29 @@ onMounted(() => {
position: relative;
}
.chat-container .conversations {
flex: 1 1 auto;
.chat-container .conversations:not(.is-open) {
width: 0;
opacity: 0;
flex: 0 0 0;
}
.chat-container .conversations.is-open {
overflow: hidden; /* 确保内容不溢出 */
white-space: nowrap; /* 防止文本换行 */
flex: 1 1 auto; /* 当侧边栏打开时,占据可用空间 */
}
.conversations {
display: flex;
flex-direction: column;
width: 100px;
width: 230px; /* 初始宽度 */
height: 100%;
max-width: 230px;
overflow-y: auto;
border-right: 1px solid var(--main-light-3);
min-width: var(--min-sider-width);
max-width: 200px;
background-color: #FAFCFD;
overflow: hidden; /* 确保内容不溢出 */
white-space: nowrap; /* 防止文本换行 */
transition: all 0.2s ease-out;
& .actions {
height: var(--header-height);
@ -234,7 +243,6 @@ onMounted(() => {
border-radius: 4px;
}
@media (max-width: 520px) {
.conversations {
position: absolute;

View File

@ -34,7 +34,12 @@
</div>
<div class="params-item">
<p>过滤低质量</p>
<a-switch v-model:checked="meta.filter" size="small" />
<a-switch v-model:checked="meta.filter" />
</div>
<div class="params-item">
<p>重写查询</p>
<a-segmented v-model:value="meta.rewrite_query" :options="['OFF', 'ON', 'HyDE']" />
</div>
</div>
</div>
@ -122,7 +127,7 @@
:auto-size="{ minRows: 2, maxRows: 10 }"
/>
<!-- :loading="state.searchLoading" -->
<a-button @click="onQuery" :disabled="queryText.length == 0">
<a-button @click="onQuery" :disabled="queryText.length == 0" :loading="state.searchLoading">
<SearchOutlined v-if="!state.searchLoading"/>检索
</a-button>
</div>
@ -183,6 +188,7 @@ const state = reactive({
const meta = reactive({
queryCount: 10,
filter: false,
rewrite_query: 'OFF',
});
const onQuery = () => {
@ -507,9 +513,8 @@ onMounted(() => {
display: flex;
flex-direction: column;
padding: 20px;
background: var(--main-light-4);
border-radius: 8px;
margin: 20px;
margin: 20px 0;
box-sizing: border-box;
gap: 12px;
@ -523,6 +528,10 @@ onMounted(() => {
margin: 0;
}
}
.params-item.col {
flex-direction: column;
}
}
}

View File

@ -3,6 +3,7 @@
<div class="setting">
<h2>设置</h2>
<h3>模型配置</h3>
<p>请在 <code>src/.env</code> 文件中配置对应的 APIKEY可参考 <code>src/.env.template</code></p>
<div class="section">
<div class="card">
<span class="label">