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 @@