diff --git a/src/config/__init__.py b/src/config/__init__.py index 655ba8b2..e40f363f 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -59,10 +59,10 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR - self.add_item("model_provider", default="zhipu", des="模型提供商", choices=list(MODEL_NAMES.keys())) - self.add_item("model_name", default=None, des="模型名称") - self.add_item("embed_model", default="zhipu-embedding-3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) - self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys())) + self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(MODEL_NAMES.keys())) + self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称") + self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) + self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys())) self.add_item("model_local_paths", default={}, des="本地模型路径") self.add_item("use_rewrite_query", default="off", des="重写查询", choices=["off", "on", "hyde"]) ### <<< 默认配置结束 diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 65c24c64..99ecb830 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -1,6 +1,5 @@ import os -from src.models.embedding import EmbeddingModel from pymilvus import MilvusClient, MilvusException from src.utils import setup_logger, hashstr logger = setup_logger("KnowledgeBase") @@ -83,7 +82,7 @@ class KnowledgeBase: def search(self, query, collection_name, limit=3): - query_vectors = self.embed_model.encode_queries([query]) + query_vectors = self.embed_model.batch_encode([query]) return self.search_by_vector(query_vectors[0], collection_name, limit) def search_by_vector(self, vector, collection_name, limit=3): diff --git a/src/core/retriever.py b/src/core/retriever.py index 7ea5454c..d86e65fc 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,4 +1,4 @@ -from src.models.embedding import Reranker +from src.models.rerank_model import get_reranker from src.utils.logging_config import setup_logger logger = setup_logger("server-common") @@ -12,7 +12,7 @@ class Retriever: self.model = model if self.config.enable_reranker: - self.reranker = Reranker(config) + self.reranker = get_reranker(config) if self.config.enable_web_search: from src.utils.web_search import WebSearcher @@ -110,15 +110,13 @@ class Retriever: for r in all_kb_res: r["file"] = kb.id2file(r["entity"]["file_id"]) - # use distance threshold to filter results - if meta.get("mode") == "search": - kb_res = all_kb_res - else: - kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] + kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] - if self.config.enable_reranker: - for r in kb_res: - r["rerank_score"] = self.reranker.compute_score([rw_query, r["entity"]["text"]], normalize=True)[0] + if self.config.enable_reranker and len(kb_res) > 0: + texts = [r["entity"]["text"] for r in kb_res] + rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True) + for i, r in enumerate(kb_res): + r["rerank_score"] = rerank_scores[i] kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) kb_res = [_res for _res in kb_res if _res["rerank_score"] > rerank_threshold] diff --git a/src/models/embedding.py b/src/models/embedding.py index 6b2f84f8..2bcd8906 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,105 +1,125 @@ import os -from FlagEmbedding import FlagModel, FlagReranker +import json +import requests +from FlagEmbedding import FlagModel -from src.config import EMBED_MODEL_INFO, RERANKER_LIST +from src.config import EMBED_MODEL_INFO from src.utils.logging_config import setup_logger from src.utils import hashstr logger = setup_logger("EmbeddingModel") -GLOBAL_EMBED_STATE = {} - - -class EmbeddingModel(FlagModel): - def __init__(self, model_info, config, **kwargs): - self.info = model_info - model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("default_path")) - logger.info(f"Loading embedding model {model_info['name']} from {model_name_or_path}") +class LocalEmbeddingModel(FlagModel): + def __init__(self, config, **kwargs): + info = EMBED_MODEL_INFO[config.embed_model] + model_name_or_path = config.model_local_paths.get(info["name"], info.get("default_path")) + logger.info(f"Loading embedding model {info['name']} from {model_name_or_path}") super().__init__(model_name_or_path, - query_instruction_for_retrieval=model_info.get("query_instruction", None), + query_instruction_for_retrieval=info.get("query_instruction", None), use_fp16=False, **kwargs) - logger.info(f"Embedding model {model_info['name']} loaded") + logger.info(f"Embedding model {info['name']} loaded") -class Reranker(FlagReranker): - def __init__(self, config, **kwargs): - - assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" - - model_info = RERANKER_LIST[config.reranker] - model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("default_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.reranker} loaded") - from zhipuai import ZhipuAI -class ZhipuEmbedding: - def __init__(self, model_info, config) -> None: - self.config = config - self.model_info = model_info - self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) - logger.info("Zhipu Embedding model loaded") - self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:" +class RemoteEmbeddingModel: + embed_state = {} - def predict(self, message): + def batch_encode(self, messages, batch_size=20): data = [] - batch_size = 20 - if len(message) > batch_size: - global GLOBAL_EMBED_STATE - task_id = hashstr(message) - logger.info(f"Creating new state for process {task_id}") - GLOBAL_EMBED_STATE[task_id] = { + if len(messages) > batch_size: + task_id = hashstr(messages) + self.embed_state[task_id] = { 'status': 'in-progress', - 'total': len(message), + 'total': len(messages), 'progress': 0 } - for i in range(0, len(message), batch_size): - if len(message) > batch_size: - logger.info(f"Encoding {i} to {i+batch_size} with {len(message)} messages") - GLOBAL_EMBED_STATE[task_id]['progress'] = i + for i in range(0, len(messages), batch_size): + group_msg = messages[i:i+batch_size] + logger.info(f"Encoding {i} to {i+batch_size} with {len(messages)} messages") + response = self.encode_queries(group_msg) + data.extend(response) - group_msg = message[i:i+batch_size] - response = self.client.embeddings.create( - model=self.model_info.get("default_path", None), - input=group_msg, - ) + if len(messages) > batch_size: + self.embed_state[task_id]['progress'] = len(messages) + self.embed_state[task_id]['status'] = 'completed' - data.extend([a.embedding for a in response.data]) + return data - if len(message) > batch_size: - GLOBAL_EMBED_STATE[task_id]['progress'] = len(message) - GLOBAL_EMBED_STATE[task_id]['status'] = 'completed' +class ZhipuEmbedding(RemoteEmbeddingModel): + def __init__(self, config) -> None: + self.config = config + self.model = EMBED_MODEL_INFO[config.embed_model]["name"] + self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) + + def predict(self, message): + response = self.client.embeddings.create( + model=self.model, + input=message, + ) + data = [a.embedding for a in response.data] return 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) +class SiliconFlowEmbedding(RemoteEmbeddingModel): + + def __init__(self, config) -> None: + self.url = "https://api.siliconflow.cn/v1/embeddings" + self.model = EMBED_MODEL_INFO[config.embed_model]["name"] + api_key = os.getenv("SILICONFLOW_API_KEY") + assert api_key, "SILICONFLOW_API_KEY is required" + self.headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json" + } + + def encode(self, message): + payload = self.build_payload(message) + response = requests.request("POST", self.url, json=payload, headers=self.headers) + response = json.loads(response.text) + # logger.debug(f"SiliconFlow Embedding response: {response}") + assert response["data"], f"SiliconFlow Embedding failed: {response}" + data = [a["embedding"] for a in response["data"]] + return data + + def encode_queries(self, queries): + return self.encode(queries) + + def build_payload(self, message): + return { + "model": self.model, + "input": message, + } + def get_embedding_model(config): if not config.enable_knowledge_base: return None + provider, model_name = config.embed_model.split('/', 1) assert config.embed_model in EMBED_MODEL_INFO.keys(), f"Unsupported embed model: {config.embed_model}, only support {EMBED_MODEL_INFO.keys()}" + logger.debug(f"Loading embedding model {config.embed_model}") + if provider == "local": + model = LocalEmbeddingModel(config) - if config.embed_model in ["bge-large-zh-v1.5"]: - model = EmbeddingModel(EMBED_MODEL_INFO[config.embed_model], config) + if provider == "zhipu": + model = ZhipuEmbedding(config) - if config.embed_model in ["zhipu-embedding-2", "zhipu-embedding-3"]: - model = ZhipuEmbedding(EMBED_MODEL_INFO[config.embed_model], config) + if provider == "siliconflow": + model = SiliconFlowEmbedding(config) return model diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py new file mode 100644 index 00000000..dbd24b0d --- /dev/null +++ b/src/models/rerank_model.py @@ -0,0 +1,70 @@ +import os +import json +import requests +import numpy as np +from FlagEmbedding import FlagReranker + +from src.config import RERANKER_LIST +from src.utils.logging_config import setup_logger + + +logger = setup_logger("RerankModel") + + +class LocalReranker(FlagReranker): + def __init__(self, config, **kwargs): + model_info = RERANKER_LIST[config.reranker] + model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("default_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.reranker} loaded") + + +def sigmoid(x): + return 1 / (1 + np.exp(-x)) + +class SilconFlowReranker(): + def __init__(self, config, **kwargs): + self.url = "https://api.siliconflow.cn/v1/rerank" + self.model = RERANKER_LIST[config.reranker]["name"] + + api_key = os.getenv("SILICONFLOW_API_KEY") + assert api_key, "SILICONFLOW_API_KEY is required" + self.headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json" + } + + def compute_score(self, sentence_pairs, batch_size = 256, max_length = 512, normalize = False): + # TODO 还没实现 batch_size + query, sentences = sentence_pairs[0], sentence_pairs[1] + payload = self.build_payload(query, sentences, max_length) + response = requests.request("POST", self.url, json=payload, headers=self.headers) + response = json.loads(response.text) + logger.debug(f"SiliconFlow Reranker response: {response}") + + results = sorted(response["results"], key=lambda x: x["index"]) + all_scores = [result["relevance_score"] for result in results] + + if normalize: + all_scores = [sigmoid(score) for score in all_scores] + + return all_scores + + def build_payload(self, query, sentences, max_length = 512): + return { + "model": self.model, + "query": query, + "documents": sentences, + "max_chunks_per_doc": max_length, + } + +def get_reranker(config): + assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" + provider, model_name = config.reranker.split('/', 1) + if provider == "local": + return LocalReranker(config) + elif provider == "siliconflow": + return SilconFlowReranker(config) + diff --git a/src/static/models.yaml b/src/static/models.yaml index ba3a6d54..ca649011 100644 --- a/src/static/models.yaml +++ b/src/static/models.yaml @@ -68,30 +68,33 @@ MODEL_NAMES: siliconflow: name: SiliconFlow url: https://cloud.siliconflow.cn/models - default: meta-llama/Meta-Llama-3.1-8B-Instruct + default: Qwen/Qwen2.5-7B-Instruct env: - SILICONFLOW_API_KEY models: - meta-llama/Meta-Llama-3.1-8B-Instruct - - meta-llama/Meta-Llama-3.1-70B-Instruct - - meta-llama/Meta-Llama-3.1-405B-Instruct + - Qwen/Qwen2.5-7B-Instruct - deepseek-ai/DeepSeek-R1 + - deepseek-ai/DeepSeek-V3 EMBED_MODEL_INFO: - bge-m3: + local/BAAI/bge-m3: name: BAAI/bge-m3 default_path: BAAI/bge-m3 dimension: 1024 - zhipu-embedding-2: - name: zhipu-embedding-2 - default_path: embedding-2 + zhipu/zhipu-embedding-2: + name: embedding-2 dimension: 1024 - zhipu-embedding-3: - name: zhipu-embedding-3 - default_path: embedding-3 + zhipu/zhipu-embedding-3: + name: embedding-3 dimension: 2048 + siliconflow/BAAI/bge-m3: + name: BAAI/bge-m3 + dimension: 1024 RERANKER_LIST: - bge-reranker-v2-m3: + local/BAAI/bge-reranker-v2-m3: name: BAAI/bge-reranker-v2-m3 default_path: BAAI/bge-reranker-v2-m3 + siliconflow/BAAI/bge-reranker-v2-m3: + name: BAAI/bge-reranker-v2-m3 diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 22fb676e..6fb4e256 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -46,7 +46,7 @@
总结对话标题
-
+
启用检索
@@ -117,13 +117,13 @@
正在检索……
-
正在思考…… {{ message.reasoning }}
+
正在思考…… {{ message.reasoning_content }}
- 请求错误,请重试 + 请求错误,请重试。{{ message.message }}
console.log(message) +const consoleMsg = (msg) => console.log(msg) onClickOutside(panel, () => setTimeout(() => opts.showPanel = false, 30)) onClickOutside(modelCard, () => setTimeout(() => opts.showModelCard = false, 30)) -const renderMarkdown = (message) => { - if (message.status === 'loading') { - return marked.parse(message.text + '🟢') +const renderMarkdown = (msg) => { + if (msg.status === 'loading') { + return marked.parse(msg.text + '🟢') } else { - return marked.parse(message.text) + return marked.parse(msg.text) } } @@ -306,11 +306,11 @@ const generateRandomHash = (length) => { return hash; } -const appendUserMessage = (message) => { +const appendUserMessage = (msg) => { conv.value.messages.push({ id: generateRandomHash(16), role: 'sent', - text: message + text: msg }) scrollToBottom() } @@ -320,7 +320,7 @@ const appendAiMessage = (text, refs=null) => { id: generateRandomHash(16), role: 'received', text: text, - reasoning: '', + reasoning_content: '', refs, status: "init", meta: {}, @@ -329,40 +329,44 @@ const appendAiMessage = (text, refs=null) => { } const updateMessage = (info) => { - const message = conv.value.messages.find((message) => message.id === info.id); - if (message) { + const msg = conv.value.messages.find((msg) => msg.id === info.id); + if (msg) { try { // 只有在 text 不为空时更新 if (info.text !== null && info.text !== undefined && info.text !== '') { - message.text = info.text; + msg.text = info.text; } - if (info.reasoning !== null && info.reasoning !== undefined && info.reasoning !== '') { - message.reasoning = info.reasoning; + if (info.reasoning_content !== null && info.reasoning_content !== undefined && info.reasoning_content !== '') { + msg.reasoning_content = info.reasoning_content; } // 只有在 refs 不为空时更新 if (info.refs !== null && info.refs !== undefined) { - message.refs = info.refs; + msg.refs = info.refs; } if (info.model_name !== null && info.model_name !== undefined && info.model_name !== '') { - message.model_name = info.model_name; + msg.model_name = info.model_name; } - // 只有在 status 不为空时更新 + // 只有在 status 不为空时更新 if (info.status !== null && info.status !== undefined && info.status !== '') { - message.status = info.status; + msg.status = info.status; } if (info.meta !== null && info.meta !== undefined) { - message.meta = info.meta; + msg.meta = info.meta; + } + + if (info.message !== null && info.message !== undefined) { + msg.message = info.message; } scrollToBottom(); } catch (error) { console.error('Error updating message:', error); - message.status = 'error'; - message.text = '消息更新失败'; + msg.status = 'error'; + msg.text = '消息更新失败'; } } else { console.error('Message not found:', info.id); @@ -371,9 +375,9 @@ const updateMessage = (info) => { const groupRefs = (id) => { - const message = conv.value.messages.find((message) => message.id === id) - if (message.refs && message.refs.knowledge_base.results.length > 0) { - message.groupedResults = message.refs.knowledge_base.results + const msg = conv.value.messages.find((msg) => msg.id === id) + if (msg.refs && msg.refs.knowledge_base.results.length > 0) { + msg.groupedResults = msg.refs.knowledge_base.results .filter(result => result.file && result.file.filename) .reduce((acc, result) => { const { filename } = result.file; @@ -387,11 +391,11 @@ const groupRefs = (id) => { scrollToBottom() } -const simpleCall = (message) => { +const simpleCall = (msg) => { return new Promise((resolve, reject) => { fetch('/api/chat/call_lite', { method: 'POST', - body: JSON.stringify({ query: message, }), + body: JSON.stringify({ query: msg, }), headers: { 'Content-Type': 'application/json' } }) .then((response) => response.json()) @@ -432,24 +436,11 @@ const fetchChatResponse = (user_input, cur_res_id) => { const readChunk = () => { return reader.read().then(({ done, value }) => { if (done) { - const message = conv.value.messages.find((message) => message.id === cur_res_id) - console.log(message) - if (message.meta.enable_retrieval) { + const msg = conv.value.messages.find((msg) => msg.id === cur_res_id) + console.log(msg) + if (msg.meta.enable_retrieval) { console.log("fetching refs") - fetchRefs(cur_res_id).then((data) => { - console.log(data) - updateMessage({ - id: cur_res_id, - refs: data, - status: "finished", - }); - groupRefs(cur_res_id); - }) - } else { - updateMessage({ - id: cur_res_id, - status: "finished", - }); + groupRefs(cur_res_id); } isStreaming.value = false; if (conv.value.messages.length === 2) { renameTitle(); } @@ -468,12 +459,11 @@ const fetchChatResponse = (user_input, cur_res_id) => { updateMessage({ id: cur_res_id, text: data.response, - reasoning: data.reasoning_response, - model_name: data.model_name, + reasoning_content: data.reasoning_content, status: data.status, meta: data.meta, + ...data, }); - // console.log(data) // console.log("Last message", conv.value.messages[conv.value.messages.length - 1].text) // console.log("Last message", conv.value.messages[conv.value.messages.length - 1].status) @@ -541,7 +531,7 @@ const sendMessage = () => { const retryMessage = (id) => { // 找到 id 对应的 message,然后删除包含 message 在内以及后面所有的 message - const index = conv.value.messages.findIndex(message => message.id === id); + const index = conv.value.messages.findIndex(msg => msg.id === id); const pastMessage = conv.value.messages[index-1] console.log("retryMessage", id, pastMessage) conv.value.inputText = pastMessage.text @@ -552,8 +542,8 @@ const retryMessage = (id) => { sendMessage(); } -const autoSend = (message) => { - conv.value.inputText = message +const autoSend = (msg) => { + conv.value.inputText = msg sendMessage() } @@ -750,12 +740,12 @@ watch( /* animation: slideInUp 0.1s ease-in; */ .err-msg { - color: #FF6B6B; - border: 1px solid #FF6B6B; - padding: 0.2rem 1rem; + color: #eb8080; + border: 1px solid #eb8080; + padding: 0.5rem 1rem; border-radius: 8px; - text-align: center; - background: #FFF0F0; + text-align: left; + background: #FFF5F5; margin-bottom: 10px; cursor: pointer; } diff --git a/web/src/components/RefsComponent.vue b/web/src/components/RefsComponent.vue index 342994f7..03fbe0f4 100644 --- a/web/src/components/RefsComponent.vue +++ b/web/src/components/RefsComponent.vue @@ -1,10 +1,10 @@