modified config embedding model and startup

This commit is contained in:
Wenjie Zhang 2024-07-22 00:00:54 +08:00
parent 590dfd1836
commit 3ab20611a2
14 changed files with 212 additions and 60 deletions

View File

@ -31,8 +31,7 @@ class Config(SimpleConfig):
def __init__(self, filename=None):
super().__init__()
self.filename = filename
logger.info(f"Loading config from {filename}")
self.filename = filename or "config/base.yaml"
### >>> 默认配置
# 可以在 config/base.yaml 中覆盖
@ -48,6 +47,8 @@ class Config(SimpleConfig):
# 模型配置
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
self.model_provider = "qianfan"
self.model_name = None # 默认使用 provider 的默认模型
self.embed_model = "bge-large-zh-v1.5"
self.reranker = "bge-reranker-v2-m3"
### <<< 默认配置结束
@ -66,6 +67,7 @@ class Config(SimpleConfig):
def load(self):
"""根据传入的文件覆盖掉默认配置"""
logger.info(f"Loading config from {self.filename}")
if self.filename is not None and os.path.exists(self.filename):
if self.filename.endswith(".json"):
with open(self.filename, 'r') as f:
@ -77,7 +79,20 @@ class Config(SimpleConfig):
logger.warning(f"Config file {self.filename} not found")
def save(self):
if self.filename is not None:
with open(self.filename, 'w') as f:
logger.info(f"Saving config to {self.filename}")
if self.filename is None:
logger.warning("Config file is not specified, save to default config/base.yaml")
self.filename = "config/base.yaml"
if self.filename.endswith(".json"):
with open(self.filename, 'w+') as f:
json.dump(self, f, indent=4, ensure_ascii=False)
elif self.filename.endswith(".yaml"):
with open(self.filename, 'w+') as f:
yaml.dump(self, f, indent=2)
else:
logger.warning(f"Unknown config file type {self.filename}, save as json")
with open(self.filename, 'w+') as f:
json.dump(self, f, indent=4)
logger.info(f"Config file {self.filename} saved")
logger.info(f"Config file {self.filename} saved")

View File

@ -6,7 +6,7 @@ from plugins import pdf2txt
from core.knowledgebase import KnowledgeBase
from core.filereader import pdfreader, plainreader
from core.graphbase import GraphDatabase
from models.embedding import EmbeddingModel
from models.embedding import get_embedding_model
logger = setup_logger("DataBaseManager")
@ -17,9 +17,10 @@ class DataBaseLite:
self.description = description
self.db_type = db_type
self.db_id = kwargs.get("db_id", hashstr(name))
self.metaname = kwargs.get("metaname", f"{db_type}_{hashstr(name)}")
self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}")
self.metadata = kwargs.get("metaname", {})
self.files = kwargs.get("files", [])
self.embed_model = kwargs.get("embed_model", None)
def update(self, metadata):
@ -31,6 +32,7 @@ class DataBaseLite:
"description": self.description,
"db_type": self.db_type,
"db_id": self.db_id,
"embed_model": self.embed_model,
"metaname": self.metaname,
"metadata": self.metadata,
"files": self.files
@ -47,7 +49,7 @@ class DataBaseManager:
def __init__(self, config=None) -> None:
self.config = config
self.database_path = "data/databases.json"
self.embed_model = EmbeddingModel(config)
self.embed_model = get_embedding_model(config)
self.knowledge_base = KnowledgeBase(config, self.embed_model)
self.data = {"databases": [], "graph": {}}
@ -96,7 +98,7 @@ class DataBaseManager:
return {"graph": {}, "message": "Graph database is not enabled"}
def create_database(self, database_name, description, db_type):
new_database = DataBaseLite(database_name, description, db_type)
new_database = DataBaseLite(database_name, description, db_type, embed_model=self.config.embed_model)
self.knowledge_base.add_collection(new_database.metaname)
self.data["databases"].append(new_database)
@ -105,6 +107,11 @@ class DataBaseManager:
def add_files(self, db_id, files):
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": "Embed model not match", "status": "failed"}
new_files = []
for file in files:
# filenames = [f["filename"] for f in db.files]
@ -138,6 +145,8 @@ class DataBaseManager:
self._save_databases()
return {"message": "全部解析完成", "status": "success"}
def get_database_info(self, db_id):
db = self.get_kb_by_id(db_id)
if db is None:

View File

@ -1,4 +1,3 @@
from core.startup import dbm, model
from models.embedding import Reranker
from utils.logging_config import setup_logger
logger = setup_logger("server-common")
@ -6,9 +5,11 @@ logger = setup_logger("server-common")
class Retriever:
def __init__(self, config):
def __init__(self, config, dbm, model):
self.config = config
self.reranker = Reranker(config)
self.dbm = dbm
self.model = model
def retrieval(self, query, history, meta):
@ -56,7 +57,7 @@ class Retriever:
results = []
if meta.get("use_graph"):
for entitie in entities:
result = dbm.graph_base.query_by_vector(entitie)
result = self.dbm.graph_base.query_by_vector(entitie)
if result != []:
results.extend(result)
return {"results": self.format_query_results(results)}
@ -65,7 +66,7 @@ class Retriever:
kb_res = []
if meta.get("db_name"):
kb_res = dbm.knowledge_base.search(query, meta["db_name"], limit=5)
kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5)
for r in kb_res:
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
@ -97,7 +98,7 @@ class Retriever:
# 构建提示词
rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query)
# 调用语言模型生成重写的查询假设使用某个API
rewritten_query = model.predict(rewritten_query_prompt).content
rewritten_query = self.model.predict(rewritten_query_prompt).content
if meta.get("use_graph"):
@ -113,7 +114,7 @@ class Retriever:
"""
# 构建提示词
entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query)
entities = model.predict(entity_extraction_prompt).content.split(",")
entities = self.model.predict(entity_extraction_prompt).content.split(",")
entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)]
else:
entities = []

View File

@ -1,10 +1,25 @@
from core import DataBaseManager
from core.retriever import Retriever
from models import select_model
from config import Config
from utils import setup_logger
logger = setup_logger("Startup")
config = Config("config/base.yaml")
model = select_model(config)
dbm = DataBaseManager(config)
class Startup:
def __init__(self):
self.config = Config("config/base.yaml")
self.model = select_model(self.config)
self.dbm = DataBaseManager(self.config)
self.retriever = Retriever(self.config, self.dbm, self.model)
# 启动本地图数据库
def restart(self):
logger.info("Restarting...")
self.model = select_model(self.config)
self.dbm = DataBaseManager(self.config)
self.retriever = Retriever(self.config, self.dbm, self.model)
logger.info("Restarted")
startup = Startup()

View File

@ -6,7 +6,7 @@ def select_model(config):
model_provider = config.model_provider
model_name = config.model_name
logger.info(f"Selecting model from {model_provider} with name {model_name}")
logger.info(f"Selecting model from {model_provider} with {model_name or 'default'}")
if model_provider == "deepseek":
from models.chat_model import DeepSeek

View File

@ -40,31 +40,39 @@ class Reranker(FlagReranker):
assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}"
model_name_or_path = config.model_local_paths.get(config.reranker, RERANKER_LIST[config.reranker])
logger.info(f"Loading Reranker model {config.re_ranker} from {model_name_or_path}")
logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}")
super().__init__(model_name_or_path, use_fp16=True, **kwargs)
logger.info(f"Reranker model {config.re_ranker} loaded")
logger.info(f"Reranker model {config.reranker} loaded")
from zhipuai import ZhipuAI
client = ZhipuAI(api_key="270ea71e9560c0ff406acbcdd48bfd97.e3XOMdWKuZb7Q1Sk")
response = client.embeddings.create(
model="embedding-2", #填写需要调用的模型名称
input=["你好","woshi"]
)
print(response.data.shape)
class ZhipuEmbedding:
def __init__(self, config) -> None:
self.config = config
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI"))
logger.info("Zhipu Embedding model loaded")
self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"
def predict(self, message):
response = self.client.embeddings.create(
model=SUPPORT_LIST[self.config.embed_model],
input=message
)
return response.data
return [a["embedding"] for a in response["data"]]
def encode(self, message):
return self.predict(message)
def encode_queries(self, queries):
# queries = [self.query_instruction_for_retrieval + query for query in queries]
return self.predict(queries)
def get_embedding_model(config):
if config.embed_model == "zhipu":
return ZhipuEmbedding(config)
else:
return EmbeddingModel(config)

View File

@ -3,13 +3,10 @@ from flask import Blueprint, jsonify, request, Response
from core import HistoryManager
from utils.logging_config import setup_logger
from core.startup import config, model
from core.retriever import Retriever
from core.startup import startup
common = Blueprint('common', __name__)
logger = setup_logger("server-common")
retriever = Retriever(config)
@common.route('/', methods=["GET"])
def route_index():
@ -35,7 +32,7 @@ def chat():
logger.debug(f"Web query: {query}")
history_manager = HistoryManager(request_data['history'])
new_query, refs = retriever(query, history_manager.messages, meta)
new_query, refs = startup.retriever(query, history_manager.messages, meta)
messages = history_manager.get_history_with_msg(new_query)
history_manager.add_user(query)
@ -43,7 +40,7 @@ def chat():
def generate_response():
content = ""
for delta in model.predict(messages, stream=True):
for delta in startup.model.predict(messages, stream=True):
if delta.content:
content += delta.content
response_chunk = json.dumps({
@ -59,7 +56,7 @@ def chat():
def call():
request_data = json.loads(request.data)
query = request_data['query']
response = model.predict(query)
response = startup.model.predict(query)
logger.debug(f"Call query: {query} Response: {response.content}")
return jsonify({
@ -68,4 +65,16 @@ def call():
@common.route('/config', methods=['get'])
def get_config():
return jsonify(config)
return jsonify(startup.config)
@common.route('/config', methods=['post'])
def update_config():
request_data = json.loads(request.data)
startup.config.update(request_data)
startup.config.save()
return jsonify(startup.config)
@common.route('/restart', methods=['POST'])
def restart():
startup.restart()
return jsonify({"message": "Restarted!"})

View File

@ -3,9 +3,8 @@ import json
import threading
from flask import Blueprint, jsonify, request, Response
from core import HistoryManager
from utils.logging_config import setup_logger
from core.startup import config, model, dbm
from core.startup import startup
db = Blueprint('database', __name__, url_prefix="/database")
@ -15,7 +14,7 @@ progress = {} # 只针对单个用户的进度
@db.route('/', methods=['GET'])
def get_databases():
database = dbm.get_databases()
database = startup.dbm.get_databases()
return jsonify(database)
@db.route('/', methods=['POST'])
@ -25,7 +24,7 @@ def create_database():
description = data.get('description')
db_type = data.get('db_type')
logger.debug(f"Create database {database_name}")
database = dbm.create_database(database_name, description, db_type)
database = startup.dbm.create_database(database_name, description, db_type)
return jsonify(database)
# TODO: 删除数据库
@ -34,7 +33,7 @@ def delete_database():
data = json.loads(request.data)
db_id = data.get('db_id')
logger.debug(f"Delete database {db_id}")
dbm.delete_database(db_id)
startup.dbm.delete_database(db_id)
return jsonify({"message": "删除成功"})
@ -44,8 +43,8 @@ def create_document_by_file():
db_id = data.get('db_id')
files = data.get('files')
logger.debug(f"Add document in {db_id} by file: {files}")
dbm.add_files(db_id, files)
return jsonify({"status": "全部解析完成"})
msg = startup.dbm.add_files(db_id, files)
return jsonify(msg)
@db.route('/info', methods=['GET'])
@ -55,7 +54,7 @@ def get_database_info():
return jsonify({"message": "db_id is required"}), 400
logger.debug(f"Get database {db_id} info")
database = dbm.get_database_info(db_id)
database = startup.dbm.get_database_info(db_id)
if database is None:
return jsonify({"message": "database not found"}), 404
@ -69,7 +68,7 @@ def delete_document():
db_id = data.get('db_id')
file_id = data.get('file_id')
logger.debug(f"DELETE document {file_id} info in {db_id}")
dbm.delete_file(db_id, file_id)
startup.dbm.delete_file(db_id, file_id)
return jsonify({"message": "删除成功"})
@db.route('/document', methods=['GET'])
@ -77,7 +76,7 @@ def get_document_info():
db_id = request.args.get('db_id')
file_id = request.args.get('file_id')
logger.debug(f"GET document {file_id} info in {db_id}")
info = dbm.get_file_info(db_id, file_id)
info = startup.dbm.get_file_info(db_id, file_id)
return jsonify(info)
@db.route('/upload', methods=['POST'])
@ -97,5 +96,5 @@ def upload_file():
@db.route('/graph', methods=['GET'])
def get_graph_info():
graph_info = dbm.get_graph()
graph_info = startup.dbm.get_graph()
return jsonify(graph_info)

View File

@ -1,5 +1,5 @@
<script setup>
import { KeepAlive } from 'vue'
import { KeepAlive, onMounted } from 'vue'
import { RouterLink, RouterView, useRoute } from 'vue-router'
import {
MessageOutlined,
@ -10,6 +10,20 @@ import {
BookFilled
} from '@ant-design/icons-vue'
import { themeConfig } from '@/assets/theme'
import { useConfigStore } from '@/stores/counter'
const configStore = useConfigStore()
const getRemoteConfig = () => {
fetch('/api/config').then(res => res.json()).then(data => {
console.log(data)
configStore.setConfig(data)
})
}
onMounted(() => {
getRemoteConfig()
})
// 使 vue3 setup composition API
const route = useRoute()
@ -72,7 +86,7 @@ div.header, #app-router-view {
flex: 0 0 80px;
justify-content: flex-start;
align-items: center;
background-color: #F4F8F9;
background-color: #F2F6F7;
height: 100%;
width: 80px;
border-right: 1px solid #e2eef3;

View File

@ -62,7 +62,7 @@ const router = createRouter({
{
path: '',
name: 'setting',
component: () => import('../views/EmptyView.vue'),
component: () => import('../views/SettingView.vue'),
meta: { keepAlive: true }
}
]

View File

@ -10,3 +10,17 @@ export const useCounterStore = defineStore('counter', () => {
return { count, doubleCount, increment }
})
export const useConfigStore = defineStore('config', () => {
const config = ref({})
function setConfig(newConfig) {
config.value = newConfig
}
function setConfigValue(key, value) {
config.value[key] = value
}
return { config, setConfig, setConfigValue }
})

View File

@ -10,10 +10,13 @@
<div class="icon"><ReadFilled /></div>
<div class="info">
<h3>{{ database.name }}</h3>
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata?.row_count }}</span></p>
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata?.row_count }}</span></p>
</div>
</div>
<p class="description">{{ database.description }}</p>
<div class="tags">
<a-tag color="blue" v-if="database.embed_model">Embed: {{ database.embed_model }}</a-tag>
</div>
</div>
<div class="sider-bottom">
</div>
@ -292,11 +295,15 @@ const addDocumentByFile = () => {
.then(data => {
console.log(data)
fileList.value = []
message.success(data.status)
if (data.status === 'failed') {
message.error(data.message)
} else {
message.success(data.message)
}
})
.catch(error => {
console.error(error)
message.error(error.status)
message.error(error.message)
})
.finally(() => {
getDatabaseInfo()

View File

@ -35,10 +35,13 @@
<div class="icon"><ReadFilled /></div>
<div class="info">
<h3>{{ database.name }}</h3>
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata.row_count }}</span></p>
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata.row_count }}</span></p>
</div>
</div>
<p class="description">{{ database.description }}</p>
<div class="tags">
<a-tag color="blue" v-if="database.embed_model">Embed: {{ database.embed_model }}</a-tag>
</div>
<!-- <button @click="deleteDatabase(database.collection_name)">删除</button> -->
</div>
</div>
@ -197,8 +200,8 @@ onMounted(() => {
.dbcard, .database {
padding: 10px;
border-radius: 12px;
width: 380px;
height: 150px;
width: 360px;
height: 160px;
padding: 20px;
cursor: pointer;
flex: 1 1 380px;
@ -226,6 +229,7 @@ onMounted(() => {
.info {
h3, p {
margin: 0;
color: black;
}
p {
@ -239,9 +243,10 @@ onMounted(() => {
color: var(--c-text-light-1);
overflow: hidden;
display: -webkit-box;
-webkit-line-clamp: 2;
-webkit-line-clamp: 1;
-webkit-box-orient: vertical;
text-overflow: ellipsis;
margin-bottom: 10px;
}
}

View File

@ -0,0 +1,56 @@
<template>
<div class="not-found">
<h1>404 - 页面还没做</h1>
<p>Sorry, Yemian has not been zuoed.</p>
<a-button @click="sendRestart" :loading="state.loading">重载</a-button>
<p>{{ configStore.config }}</p>
</div>
</template>
<script setup>
import { message } from 'ant-design-vue';
import { reactive, ref } from 'vue'
import { useConfigStore } from '@/stores/counter';
const configStore = useConfigStore()
const state = reactive({
loading: false,
})
const sendRestart = () => {
console.log('Restarting...')
state.loading = true
fetch('/api/restart', {
method: 'POST',
}).then(() => {
console.log('Restarted')
state.loading = false
message.success('重载成功')
setTimeout(() => {
window.location.reload()
}, 1000)
})
}
</script>
<style scoped>
.not-found {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
height: 80vh;
text-align: center;
}
.not-found h1 {
font-size: 2rem;
margin-bottom: 1rem;
}
.not-found p {
font-size: 1.5rem;
margin-bottom: 2rem;
}
</style>