From d9e39a7a7dd33fbaf9d14156f522af4a6e8e71e8 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 20 Feb 2025 01:26:12 +0800 Subject: [PATCH] Merge branch 'main' of https://github.com/xerrors/Yuxi-Know --- .gitignore | 4 +- README.md | 31 +-- docker/api.Dockerfile | 7 +- docker/docker-compose.dev.yml | 6 +- requirements.txt | 5 +- scripts/init.sh | 25 --- src/config/__init__.py | 26 ++- src/core/retriever.py | 34 +++- src/core/startup.py | 5 + src/models/__init__.py | 6 +- src/models/chat_model.py | 4 +- src/models/embedding.py | 16 +- src/models/ollama_embedding.py | 179 +++++++++++++++++ src/plugins/oneke.py | 4 +- src/routers/chat_router.py | 45 ++++- src/static/config.dev.yaml | 0 src/static/config.yaml | 0 src/{config => static}/models.yaml | 13 +- src/utils/web_search.py | 69 +++++++ test/test_concurrency.py | 121 +++++++++++ web/src/components/ChatComponent.vue | 92 ++++++--- web/src/components/RefsComponent.vue | 4 +- web/src/components/TableConfigComponent.vue | 212 ++++++++++++++++++++ web/src/router/index.js | 2 +- web/src/views/DataBaseInfoView.vue | 11 +- web/src/views/SettingView.vue | 38 +++- 26 files changed, 809 insertions(+), 150 deletions(-) delete mode 100644 scripts/init.sh create mode 100644 src/models/ollama_embedding.py create mode 100644 src/static/config.dev.yaml create mode 100644 src/static/config.yaml rename src/{config => static}/models.yaml (90%) create mode 100644 src/utils/web_search.py create mode 100644 test/test_concurrency.py create mode 100644 web/src/components/TableConfigComponent.vue diff --git a/.gitignore b/.gitignore index df19050c..ad909208 100644 --- a/.gitignore +++ b/.gitignore @@ -33,6 +33,8 @@ src/data */package-lock.json web/package-lock.json saves +saves_dev notebooks graphrag -docker/volumes \ No newline at end of file +docker/volumes +.cursorrules diff --git a/README.md b/README.md index 17daf2ac..38cfd543 100644 --- a/README.md +++ b/README.md @@ -12,9 +12,8 @@ > [!NOTE] > 当前项目还处于开发的早期,还存在一些 BUG,有问题随时提 issue。 -已知问题: - -- [ ] 从 Flask 更换到 Fast API 之后,并行命令还存在问题。 +- [ ] Ollma Embedding 支持(Open-like Embedding 支持) +- [x] DeepSeek-R1 支持,需配置 `DEEPSEEK_API_KEY` 或者 `SILICONFLOW_API_KEY` 使用 ## 概述 @@ -31,12 +30,12 @@ ZHIPUAI_API_KEY=270ea********8bfa97.e3XOMd****Q1Sk OPENAI_API_KEY=sk-*********[可选] ``` -本项目的基础对话服务可以在不含显卡的设备上运行,大模型使用在线服务商的接口。但是如果想要完整的知识库对话体验,则需要 8G 以上的显存。因为需要本地运行 embedding 模型和 rerank 模型。 +本项目的基础对话服务可以在不含显卡的设备上运行,大模型使用在线服务商的接口。但是如果想要完整的知识库对话体验,则需要 8G 以上的显存。因为需要本地运行 embedding 模型和 rerank 模型。如果需要指定本地模型所在路径,需要配置 `MODEL_DIR` 参数。 **提醒**:下面的脚本会启动开发版本,源代码的修改会自动更新(含前端和后端)。如果生产环境部署,请使用 `docker/docker-compose.yml` 启动。 ```bash -docker-compose -f docker/docker-compose.dev.yml up --build +docker compose -f docker/docker-compose.dev.yml --env-file src/.env up --build ``` **也可以加上 `-d` 参数,后台运行。* @@ -63,7 +62,7 @@ docker-compose -f docker/docker-compose.dev.yml up --build 关闭 docker 服务: ```bash -docker-compose -f docker/docker-compose.dev.yml down +docker compose -f docker/docker-compose.dev.yml --env-file src/.env down ``` 查看日志: @@ -72,26 +71,10 @@ docker-compose -f docker/docker-compose.dev.yml down docker logs # 例如:docker logs api-dev ``` -如果需要使用到本地模型(不推荐手动指定),比如向量模型或者重排序模型,则需要将环境变量中设置的 `MODEL_ROOT_DIR` 做映射,比如本地模型都是存放在 `/hdd/models` 里面,则需要在 `docker-compose.yml` 和 `docker-compose.dev.yml` 中添加: - -```yml -services: - api: - build: - context: .. - dockerfile: docker/api.Dockerfile - container_name: api-dev - working_dir: /app - volumes: - - ../src:/app/src - - ../saves:/app/saves - - /hdd/zwj/models:/hdd/zwj/models # <== 修改这一行 -``` - **生产环境部署**:本项目同时支持使用 Docker 部署生产环境,只需要更换 `docker-compose` 文件就可以了。 ```bash -docker-compose -f docker/docker-compose.yml up --build +docker compose -f docker/docker-compose.yml --env-file src/.env up --build ``` ## 模型支持 @@ -131,7 +114,7 @@ docker-compose -f docker/docker-compose.yml up --build 对于**语言模型**,并不支持直接运行本地语言模型,请使用 vllm 或者 ollama 转成 API 服务之后使用。 -对于**向量模型**和**重排序模型**,可以不做修改会自动下载模型,如果下载过程中出现问题,请参考 [HF-Mirror](https://hf-mirror.com/) 配置相关内容。如果想要使用本地已经下载好的模型(不建议),可以在 `saves/config/config.yaml` 配置相关内容。同时注意要在 docker 中做映射,参考 README 中的 `docker/docker-compose.yml`。 +对于**向量模型**和**重排序模型**,可以不做修改会自动下载模型,如果下载过程中出现问题,请参考 [HF-Mirror](https://hf-mirror.com/) 配置相关内容。如果想要使用本地已经下载好的模型(不建议),可以在网页的 settings 里面做映射。 例如: diff --git a/docker/api.Dockerfile b/docker/api.Dockerfile index 31b41002..1bf10287 100644 --- a/docker/api.Dockerfile +++ b/docker/api.Dockerfile @@ -10,7 +10,12 @@ COPY ../requirements.txt /app/requirements.txt # 安装依赖(Docker 会缓存这一步,除非 requirements.txt 发生变化) RUN pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple -RUN pip install gunicorn +RUN pip install -U gunicorn -i https://pypi.tuna.tsinghua.edu.cn/simple + +RUN sed -i s@/archive.ubuntu.com/@/mirrors.aliyun.com/@g /etc/apt/sources.list +RUN sed -i s@/security.ubuntu.com/@/mirrors.aliyun.com/@g /etc/apt/sources.list +RUN apt-get clean +RUN apt-get update && apt-get install ffmpeg libsm6 libxext6 -y # 复制代码到容器中 COPY ../src /app/src diff --git a/docker/docker-compose.dev.yml b/docker/docker-compose.dev.yml index 740985b6..3d338048 100644 --- a/docker/docker-compose.dev.yml +++ b/docker/docker-compose.dev.yml @@ -7,8 +7,9 @@ services: working_dir: /app volumes: - ../src:/app/src - - ../saves:/app/saves - - /hdd/zwj/models:/hdd/zwj/models + - ../saves_dev:/app/saves + - ../src/static/config.dev.yaml:/app/src/static/config.yaml + - ${MODEL_DIR}:${MODEL_DIR} ports: - "5000:5000" depends_on: @@ -21,6 +22,7 @@ services: - NEO4J_USERNAME=neo4j - NEO4J_PASSWORD=0123456789 - MILVUS_URI=http://milvus:19530 + - MODEL_DIR=${MODEL_DIR} # 优先级高于 .env 中的 MODEL_DIR command: uvicorn src.main:app --host 0.0.0.0 --port 5000 --reload web: diff --git a/requirements.txt b/requirements.txt index ed96727a..ca3b2224 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,7 @@ dashscope>=1.20.3 PyMuPDF==1.23.26 peft>=0.11.1 -FlagEmbedding>=1.2.11 +FlagEmbedding>=1.3.2 Flask>=3.0.3 Flask_Cors>=4.0.1 llama_index>=0.11.8 @@ -24,4 +24,5 @@ opencv-python-headless docx2txt uvicorn[standard] fastapi -python-multipart \ No newline at end of file +python-multipart +tavily-python \ No newline at end of file diff --git a/scripts/init.sh b/scripts/init.sh deleted file mode 100644 index 9f6b6881..00000000 --- a/scripts/init.sh +++ /dev/null @@ -1,25 +0,0 @@ -#!/bin/bash - -# 检查是否提供了 API_KEY 参数 -if [ -z "$1" ]; then - echo "请提供 API_KEY。" - exit 1 -fi - -# 获取当前目录路径 -CURRENT_DIR=$(pwd) - -# 如果 src 目录不存在则创建 -if [! -d "${CURRENT_DIR}/src" ]; then - mkdir -p "${CURRENT_DIR}/src" -fi - -# 如果.env 文件不存在,则从.env.template 复制一份创建 -if [! -f "${CURRENT_DIR}/src/.env" ]; then - cp "${CURRENT_DIR}/src/.env.template" "${CURRENT_DIR}/src/.env" -fi - -# 将 API_KEY 写入.env 文件 -echo "ZHIPUAI_API_KEY=$1" > "${CURRENT_DIR}/src/.env" - -echo "API_KEY 已成功写入 src/.env 文件。" \ No newline at end of file diff --git a/src/config/__init__.py b/src/config/__init__.py index 5be0e905..9f5c32b5 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -6,7 +6,7 @@ from src.utils.logging_config import setup_logger logger = setup_logger("Config") -with open(Path("src/config/models.yaml"), "r") as f: +with open(Path("src/static/models.yaml"), "r") as f: _models = yaml.safe_load(f) MODEL_NAMES = _models["MODEL_NAMES"] @@ -17,7 +17,7 @@ RERANKER_LIST = _models["RERANKER_LIST"] class SimpleConfig(dict): def __key(self, key): - return "" if key is None else key.lower() + return "" if key is None else key.lower() # 目前忘记了这里为什么要 lower 了,只能说配置项最好不要有大写的 def __str__(self): return json.dumps(self) @@ -40,32 +40,33 @@ class SimpleConfig(dict): class Config(SimpleConfig): - def __init__(self, filename=None): + def __init__(self): super().__init__() self._config_items = {} + self.save_dir = "saves" + self.filename = str(Path("src/static/config.yaml")) + os.makedirs(os.path.dirname(self.filename), exist_ok=True) ### >>> 默认配置 # 可以在 config/base.yaml 中覆盖 self.add_item("stream", default=True, des="是否开启流式输出") - self.add_item("save_dir", default="saves", des="保存目录") # 功能选项 self.add_item("enable_reranker", default=False, des="是否开启重排序") self.add_item("enable_knowledge_base", default=False, des="是否开启知识库") self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱") self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎") - + self.add_item("enable_web_search", default=False, des="是否开启网页搜索") # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 - ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths + ## 如果需要自定义本地模型路径,则在 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_local_paths", default={}, des="本地模型路径") + self.add_item("use_rewrite_query", default="off", des="重写查询", choices=["off", "on", "hyde"]) ### <<< 默认配置结束 - self.filename = filename or os.path.join(self.save_dir, "config", "config.yaml") - os.makedirs(os.path.dirname(self.filename), exist_ok=True) self.load() self.handle_self() @@ -89,7 +90,9 @@ class Config(SimpleConfig): def handle_self(self): self.model_names = MODEL_NAMES model_provider_info = self.model_names.get(self.model_provider, {}) + self.model_dir = os.environ.get("MODEL_DIR", "") + # 检查模型提供商是否存在 if self.model_provider != "custom": if self.model_name not in model_provider_info["models"]: logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name") @@ -103,6 +106,7 @@ class Config(SimpleConfig): logger.warning(f"Model name {self.model_name} not in custom models, using default model name") self.model_name = self.custom_models[0]["custom_id"] + # 检查模型提供商的环境变量 conds = {} self.model_provider_status = {} for provider in self.model_names: @@ -110,10 +114,14 @@ class Config(SimpleConfig): conds_bool = [bool(os.getenv(_k)) for _k in conds[provider]] self.model_provider_status[provider] = all(conds_bool) + # 检查web_search的环境变量 + if self.enable_web_search and not os.getenv("TAVILY_API_KEY"): + logger.warning("TAVILY_API_KEY not set, web search will be disabled") + self.enable_web_search = False + self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v] assert len(self.valuable_model_provider) > 0, f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}" - def load(self): """根据传入的文件覆盖掉默认配置""" diff --git a/src/core/retriever.py b/src/core/retriever.py index 3d21971e..7ea5454c 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -14,6 +14,10 @@ class Retriever: if self.config.enable_reranker: self.reranker = Reranker(config) + if self.config.enable_web_search: + from src.utils.web_search import WebSearcher + self.web_searcher = WebSearcher() + def retrieval(self, query, history, meta): refs = {"query": query, "history": history, "meta": meta} @@ -21,6 +25,7 @@ class Retriever: 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) + refs["web_search"] = self.query_web(query, history, refs) return refs @@ -44,6 +49,12 @@ class Retriever: ) external_parts.extend(["图数据库信息:", db_text]) + # 解析网络搜索的结果 + web_res = refs.get("web_search", {}).get("results", []) + if web_res: + web_text = "\n".join(f"{r['title']}: {r['snippet']}" for r in web_res) + external_parts.extend(["网络搜索信息:", web_text]) + # 构造查询 from src.utils.prompts import knowbase_qa_template if external_parts and len(external_parts) > 0: @@ -60,8 +71,6 @@ class Retriever: raise NotImplementedError def query_graph(self, query, history, refs): - # res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content - results = [] if refs["meta"].get("use_graph") and self.config.enable_knowledge_base: for entity in refs["entities"]: @@ -70,6 +79,7 @@ class Retriever: results.extend(result) return {"results": self.format_query_results(results)} + def query_knowledgebase(self, query, history, refs): """查询知识库""" @@ -116,9 +126,27 @@ class Retriever: return {"results": kb_res, "all_results": all_kb_res, "rw_query": rw_query} + def query_web(self, query, history, refs): + """查询网络""" + + if not (refs["meta"].get("enable_web_search") and self.config.enable_web_search): + return {"results": [], "message": "Web search is disabled"} + + try: + search_results = self.web_searcher.search(query) + except Exception as e: + logger.error(f"Web search error: {str(e)}") + return {"results": [], "message": "Web search error"} + + return {"results": search_results} + def rewrite_query(self, query, history, refs): """重写查询""" - rewrite_query_span = refs["meta"].get("rewriteQuery", "off") + if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置 + rewrite_query_span = refs["meta"].get("use_rewrite_query", "off") + else: + rewrite_query_span = self.config.use_rewrite_query + if rewrite_query_span == "off": rewritten_query = query else: diff --git a/src/core/startup.py b/src/core/startup.py index 943dd10b..af1ba422 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -1,3 +1,4 @@ +import os from src.core import DataBaseManager from src.core.retriever import Retriever from src.models import select_model @@ -17,6 +18,10 @@ class Startup: self.dbm = DataBaseManager(self.config) self.retriever = Retriever(self.config, self.dbm, self.model) + self.model_lite = select_model(self.config, + model_provider=self.config.model_provider_lite or "zhipu", + model_name=self.config.model_name_lite or "glm-4-flash") + def restart(self): logger.info("Restarting...") self.start() diff --git a/src/models/__init__.py b/src/models/__init__.py index b5efa9d0..183bf358 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -1,10 +1,10 @@ from src.utils.logging_config import logger -def select_model(config): +def select_model(config, model_provider=None, model_name=None): - model_provider = config.model_provider - model_name = config.model_name + model_provider = model_provider or config.model_provider + model_name = model_name or config.model_name logger.info(f"Selecting model from {model_provider} with {model_name}") diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 8f1924e7..2c160758 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -50,8 +50,8 @@ class OpenModel(OpenAIBase): class DeepSeek(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "deepseek-chat" - api_key = os.getenv("DEEPSEEK_API_KEY") - base_url = "https://api.deepseek.com" + api_key = os.getenv("DEEPSEEK_API_KEY", "your-default-api-key") + base_url = os.getenv("DEEPSEEK_API_BASE", "https://api.deepseek.com/v1") super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) diff --git a/src/models/embedding.py b/src/models/embedding.py index 7b576abe..6b2f84f8 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,5 +1,4 @@ import os -import uuid from FlagEmbedding import FlagModel, FlagReranker from src.config import EMBED_MODEL_INFO, RERANKER_LIST @@ -15,11 +14,7 @@ GLOBAL_EMBED_STATE = {} class EmbeddingModel(FlagModel): def __init__(self, model_info, config, **kwargs): self.info = model_info - model_name_or_path = handle_local_model( - paths=config.model_local_paths, - model_name=model_info["name"], - default_path=model_info.get("default_path", None)) - + 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}") super().__init__(model_name_or_path, @@ -34,11 +29,8 @@ class Reranker(FlagReranker): assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" - model_name_or_path = handle_local_model( - paths=config.model_local_paths, - model_name=config.reranker, - default_path=RERANKER_LIST[config.reranker]) - + 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) @@ -113,6 +105,4 @@ def get_embedding_model(config): def handle_local_model(paths, model_name, default_path): model_path = paths.get(model_name, default_path) - if os.getenv("MODEL_ROOT_DIR") and not os.path.isabs(model_path): - model_path = os.path.join(os.getenv("MODEL_ROOT_DIR"), model_path) return model_path \ No newline at end of file diff --git a/src/models/ollama_embedding.py b/src/models/ollama_embedding.py new file mode 100644 index 00000000..1b099be9 --- /dev/null +++ b/src/models/ollama_embedding.py @@ -0,0 +1,179 @@ +import os +import requests +import numpy as np +from typing import List, Union, Dict +from src.utils.logging_config import setup_logger + +logger = setup_logger("OllamaEmbedding") + +class OllamaEmbedding: + """ + 使用 Ollama API 进行文本嵌入的类 + """ + def __init__(self, model_info: Dict, config) -> None: + """ + 初始化 Ollama Embedding 模型 + + Args: + model_info: 模型信息字典 + config: 配置对象 + """ + self.config = config + self.model_info = model_info + self.base_url = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434") + self.model_name = model_info.get("name", "nomic-embed-text") + self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:" + logger.info(f"Ollama Embedding model {self.model_name} initialized") + + def _get_embedding(self, text: str) -> List[float]: + """ + 获取单个文本的嵌入向量 + + Args: + text: 输入文本 + + Returns: + 嵌入向量 + """ + url = f"{self.base_url}/api/embeddings" + try: + response = requests.post(url, json={ + "model": self.model_name, + "prompt": text + }) + response.raise_for_status() + return response.json()["embedding"] + except Exception as e: + logger.error(f"Error getting embedding: {str(e)}") + raise + + def predict(self, messages: List[str]) -> List[List[float]]: + """ + 批量获取文本嵌入向量 + + Args: + messages: 文本列表 + + Returns: + 嵌入向量列表 + """ + embeddings = [] + batch_size = 20 + + for i in range(0, len(messages), batch_size): + batch = messages[i:i + batch_size] + logger.info(f"Processing batch {i//batch_size + 1}, size: {len(batch)}") + + batch_embeddings = [] + for text in batch: + embedding = self._get_embedding(text) + batch_embeddings.append(embedding) + + embeddings.extend(batch_embeddings) + + return embeddings + + def encode(self, messages: Union[str, List[str]]) -> List[List[float]]: + """ + 编码文本 + + Args: + messages: 单个文本或文本列表 + + Returns: + 嵌入向量列表 + """ + if isinstance(messages, str): + messages = [messages] + return self.predict(messages) + + def encode_queries(self, queries: List[str]) -> List[List[float]]: + """ + 编码查询文本 + + Args: + queries: 查询文本列表 + + Returns: + 查询文本的嵌入向量列表 + """ + return self.predict(queries) + + +class OllamaReranker: + """ + 使用 Ollama API 进行文本重排序的类 + """ + def __init__(self, config) -> None: + """ + 初始化 Ollama Reranker + + Args: + config: 配置对象 + """ + self.config = config + self.base_url = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434") + self.model_name = config.reranker + logger.info(f"Ollama Reranker model {self.model_name} initialized") + + def compute_score(self, query: str, passage: str) -> float: + """ + 计算查询和文本段落之间的相关性分数 + + Args: + query: 查询文本 + passage: 段落文本 + + Returns: + 相关性分数 + """ + prompt = f"Query: {query}\nPassage: {passage}\nRate the relevance of the passage to the query on a scale of 0 to 1:" + + try: + response = requests.post( + f"{self.base_url}/api/generate", + json={ + "model": self.model_name, + "prompt": prompt, + "stream": False + } + ) + response.raise_for_status() + + # 提取生成的数字作为分数 + result = response.json()["response"].strip() + try: + score = float(result) + return min(max(score, 0), 1) # 确保分数在 0-1 之间 + except ValueError: + logger.warning(f"Could not parse score from response: {result}") + return 0.0 + + except Exception as e: + logger.error(f"Error computing rerank score: {str(e)}") + return 0.0 + + def rerank(self, query: str, passages: List[str], top_n: int = None) -> List[Dict]: + """ + 重新排序文本段落 + + Args: + query: 查询文本 + passages: 段落文本列表 + top_n: 返回前 n 个结果 + + Returns: + 排序后的结果列表,每个元素包含索引和分数 + """ + scores = [] + for i, passage in enumerate(passages): + score = self.compute_score(query, passage) + scores.append({"index": i, "score": score}) + + # 按分数降序排序 + sorted_results = sorted(scores, key=lambda x: x["score"], reverse=True) + + if top_n: + sorted_results = sorted_results[:top_n] + + return sorted_results \ No newline at end of file diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index aea4c5c8..3ca3d60b 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -16,8 +16,6 @@ logger = setup_logger("OneKE") dotenv.load_dotenv() -MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR', './'), 'OneKE') - instruction_mapper = { 'NERzh': "你是专门进行实体抽取的专家。请从input中抽取出符合schema定义的实体,不存在的实体类型返回空列表。请按照JSON字符串的格式回答。", 'REzh': "你是专门进行关系抽取的专家。请从input中抽取出符合schema定义的关系三元组。请按照JSON字符串的格式回答。", @@ -42,7 +40,7 @@ class OneKE: def __init__(self, config=None): self.config = config - model_name_or_path = config.model_local_paths.get('oneke', "zjunlp/OneKE") + model_name_or_path = config.model_local_paths.get('zjunlp/OneKE', "zjunlp/OneKE") logger.info(f"Loading KGC model OneKE from {model_name_or_path}") model_config = AutoConfig.from_pretrained(model_name_or_path, trust_remote_code=True) diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index 0f5659ef..711fd1f8 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -27,9 +27,10 @@ def chat_post( history_manager = HistoryManager(history) - def make_chunk(content, status, history): + def make_chunk(content=None, status=None, history=None, reasoning_content=None): return json.dumps({ "response": content, + "reasoning_response": reasoning_content, "history": history, "model_name": startup.config.model_name, "status": status, @@ -37,33 +38,44 @@ def chat_post( }, ensure_ascii=False).encode('utf-8') + b"\n" def generate_response(): + modified_query = query - if meta.get("enable_retrieval"): - chunk = make_chunk("", "searching", history=None) + # 处理知识库检索 + if meta and meta.get("enable_retrieval"): + chunk = make_chunk(status="searching") yield chunk - new_query, refs = startup.retriever(query, history_manager.messages, meta) + modified_query, refs = startup.retriever(modified_query, history_manager.messages, meta) refs_pool[cur_res_id] = refs - else: - new_query = query - messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round')) - history_manager.add_user(query) + messages = history_manager.get_history_with_msg(modified_query, max_rounds=meta.get('history_round')) + history_manager.add_user(query) # 注意这里使用原始查询 logger.debug(f"Web history: {history_manager.messages}") content = "" + reasoning_content = "" for delta in startup.model.predict(messages, stream=True): - if not delta.content: + if not delta.content and hasattr(delta, 'reasoning_content'): + reasoning_content += delta.reasoning_content + chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning") + yield chunk continue + # 文心一言 if hasattr(delta, 'is_full') and delta.is_full: content = delta.content else: - content += delta.content + content += delta.content or "" - chunk = make_chunk(content, "loading", history=history_manager.update_ai(content)) + chunk = make_chunk(content=content, + reasoning_content=reasoning_content, + status="loading", + history=history_manager.update_ai(content)) yield chunk + logger.debug(f"Final response: {content}") + logger.debug(f"Final reasoning response: {reasoning_content}") + return StreamingResponse(generate_response(), media_type='application/json') @chat.post("/call") @@ -77,6 +89,17 @@ async def call(query: str = Body(...), meta: dict = Body(None)): return {"response": response.content} +@chat.post("/call_lite") +async def call(query: str = Body(...), meta: dict = Body(None)): + async def predict_async(query): + loop = asyncio.get_event_loop() + return await loop.run_in_executor(executor, startup.model_lite.predict, query) + + response = await predict_async(query) + logger.debug({"query": query, "response": response.content}) + + return {"response": response.content} + @chat.get("/refs") def get_refs(cur_res_id: str): global refs_pool diff --git a/src/static/config.dev.yaml b/src/static/config.dev.yaml new file mode 100644 index 00000000..e69de29b diff --git a/src/static/config.yaml b/src/static/config.yaml new file mode 100644 index 00000000..e69de29b diff --git a/src/config/models.yaml b/src/static/models.yaml similarity index 90% rename from src/config/models.yaml rename to src/static/models.yaml index f42781c2..ba3a6d54 100644 --- a/src/config/models.yaml +++ b/src/static/models.yaml @@ -18,6 +18,7 @@ MODEL_NAMES: - DEEPSEEK_API_KEY models: - deepseek-chat + - deepseek-reasoner zhipu: name: 智谱AI (Zhipu) url: https://open.bigmodel.cn/dev/api @@ -74,13 +75,13 @@ MODEL_NAMES: - meta-llama/Meta-Llama-3.1-8B-Instruct - meta-llama/Meta-Llama-3.1-70B-Instruct - meta-llama/Meta-Llama-3.1-405B-Instruct + - deepseek-ai/DeepSeek-R1 EMBED_MODEL_INFO: - bge-large-zh-v1.5: - name: bge-large-zh-v1.5 - default_path: BAAI/bge-large-zh-v1.5 + bge-m3: + name: BAAI/bge-m3 + default_path: BAAI/bge-m3 dimension: 1024 - query_instruction: "为这个句子生成表示以用于检索相关文章:" zhipu-embedding-2: name: zhipu-embedding-2 default_path: embedding-2 @@ -91,4 +92,6 @@ EMBED_MODEL_INFO: dimension: 2048 RERANKER_LIST: - bge-reranker-v2-m3: BAAI/bge-reranker-v2-m3 \ No newline at end of file + bge-reranker-v2-m3: + name: BAAI/bge-reranker-v2-m3 + default_path: BAAI/bge-reranker-v2-m3 diff --git a/src/utils/web_search.py b/src/utils/web_search.py new file mode 100644 index 00000000..781e4553 --- /dev/null +++ b/src/utils/web_search.py @@ -0,0 +1,69 @@ +import os +from typing import List, Dict +from tavily import TavilyClient +from src.utils.logging_config import setup_logger + +logger = setup_logger("web-search") + +class WebSearcher: + def __init__(self): + api_key = os.getenv("TAVILY_API_KEY") + if not api_key: + raise ValueError("TAVILY_API_KEY environment variable is not set") + self.client = TavilyClient(api_key) + logger.info("WebSearcher initialized with Tavily client") + + def search(self, query: str, max_results: int = 1) -> List[Dict]: + """ + 使用 Tavily 搜索相关内容 + + Args: + query: 搜索查询 + max_results: 最大返回结果数 + + Returns: + 搜索结果列表 + """ + try: + search_results = self.client.search( + query=query, + search_depth="basic", + max_results=max_results + ) + + # 提取需要的信息 + formatted_results = [] + for result in search_results['results'][:max_results]: + formatted_results.append({ + 'title': result.get('title', ''), + 'content': result.get('content', ''), + 'url': result.get('url', ''), + 'score': result.get('score', 0) + }) + + return formatted_results + + except Exception as e: + logger.error(f"Error during web search: {str(e)}") + return [] + + def format_search_results(self, results: List[Dict]) -> str: + """ + 将搜索结果格式化为文本 + + Args: + results: 搜索结果列表 + + Returns: + 格式化后的文本 + """ + if not results: + return "没有找到相关的网络搜索结果。" + + formatted_text = "以下是相关的网络搜索结果:\n\n" + for i, result in enumerate(results, 1): + formatted_text += f"{i}. {result['title']}\n" + formatted_text += f" {result['content']}\n" + formatted_text += f" 来源: {result['url']}\n\n" + + return formatted_text \ No newline at end of file diff --git a/test/test_concurrency.py b/test/test_concurrency.py new file mode 100644 index 00000000..2c4c4b08 --- /dev/null +++ b/test/test_concurrency.py @@ -0,0 +1,121 @@ +import asyncio +import aiohttp +import time +import json +from typing import List + +async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict: + """发送单个请求到API""" + url = "http://localhost:5000/chat/call" + payload = { + "query": "写一个冒泡排序", + "meta": {} + } + + start_time = time.time() + print(f"请求 {request_id} 开始时间: {time.strftime('%H:%M:%S', time.localtime(start_time))}") + try: + async with session.post(url, json=payload) as response: + result = await response.json() + end_time = time.time() + duration = end_time - start_time + print(f"请求 {request_id} 完成时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)") + return { + "request_id": request_id, + "status": response.status, + "time": duration, + "start_time": start_time, + "end_time": end_time, + "success": True + } + except Exception as e: + end_time = time.time() + duration = end_time - start_time + print(f"请求 {request_id} 失败时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)") + return { + "request_id": request_id, + "status": None, + "time": duration, + "start_time": start_time, + "end_time": end_time, + "success": False, + "error": str(e) + } + +async def run_concurrent_test(num_requests: int = 10) -> List[dict]: + """运行并发测试""" + async with aiohttp.ClientSession() as session: + tasks = [make_request(session, i) for i in range(num_requests)] + return await asyncio.gather(*tasks) + +def analyze_results(results: List[dict]) -> None: + """分析并打印测试结果""" + total_requests = len(results) + successful_requests = sum(1 for r in results if r["success"]) + failed_requests = total_requests - successful_requests + + response_times = [r["time"] for r in results if r["success"]] + if response_times: + avg_time = sum(response_times) / len(response_times) + max_time = max(response_times) + min_time = min(response_times) + else: + avg_time = max_time = min_time = 0 + + print(f"\n=== 并发测试结果 ===") + print(f"总请求数: {total_requests}") + print(f"成功请求: {successful_requests}") + print(f"失败请求: {failed_requests}") + print(f"平均响应时间: {avg_time:.2f} 秒") + print(f"最长响应时间: {max_time:.2f} 秒") + print(f"最短响应时间: {min_time:.2f} 秒") + + if failed_requests > 0: + print("\n失败的请求:") + for result in results: + if not result["success"]: + print(f"请求 ID {result['request_id']}: {result.get('error', '未知错误')}") + + # 添加请求时间线分析 + print("\n=== 请求时间线分析 ===") + sorted_results = sorted(results, key=lambda x: x["start_time"]) + test_start_time = sorted_results[0]["start_time"] + + print("\n时间线详情:") + print("请求ID 开始时间 结束时间 耗时(秒) 重叠请求数") + active_requests = [] + + for result in sorted_results: + # 计算当前时间点的活跃请求数 + start_time = result["start_time"] + end_time = result["end_time"] + + # 清理已完成的请求 + active_requests = [t for t in active_requests if t > start_time] + active_requests.append(end_time) + + print(f"{result['request_id']:^7} {time.strftime('%H:%M:%S', time.localtime(start_time))} " + f"{time.strftime('%H:%M:%S', time.localtime(end_time))} " + f"{result['time']:^8.2f} {len(active_requests):^6}") + + # 计算最大并发数 + max_concurrent = 0 + timeline = [] + for r in results: + timeline.append((r["start_time"], 1)) + timeline.append((r["end_time"], -1)) + + timeline.sort(key=lambda x: x[0]) + current_concurrent = 0 + for _, change in timeline: + current_concurrent += change + max_concurrent = max(max_concurrent, current_concurrent) + + print(f"\n最大并发请求数: {max_concurrent}") + +if __name__ == "__main__": + NUM_REQUESTS = 100 # 设置并发请求数 + + print(f"开始运行 {NUM_REQUESTS} 个并发请求的测试...") + results = asyncio.run(run_concurrent_test(NUM_REQUESTS)) + analyze_results(results) diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 13b37d58..768abd0b 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -81,9 +81,12 @@
搜索引擎(Bing)
-
- 重写查询 +
+ 网页搜索
+
@@ -114,6 +117,7 @@
正在检索……
+
正在思考…… {{ message.reasoning }}
{ scrollToBottom() } -const appendAiMessage = (message, refs=null) => { +const appendAiMessage = (text, refs=null) => { conv.value.messages.push({ id: generateRandomHash(16), role: 'received', - text: message, + text: text, + reasoning: '', refs, status: "init", meta: {}, @@ -329,33 +331,42 @@ const updateMessage = (info) => { const message = conv.value.messages.find((message) => message.id === info.id); if (message) { - // 只有在 text 不为空时更新 - if (info.text !== null && info.text !== undefined && info.text !== '') { - message.text = info.text; - } + try { + // 只有在 text 不为空时更新 + if (info.text !== null && info.text !== undefined && info.text !== '') { + message.text = info.text; + } - // 只有在 refs 不为空时更新 - if (info.refs !== null && info.refs !== undefined) { - message.refs = info.refs; - } + if (info.reasoning !== null && info.reasoning !== undefined && info.reasoning !== '') { + message.reasoning = info.reasoning; + } - if (info.model_name !== null && info.model_name !== undefined && info.model_name !== '') { - message.model_name = info.model_name; - } + // 只有在 refs 不为空时更新 + if (info.refs !== null && info.refs !== undefined) { + message.refs = info.refs; + } + + if (info.model_name !== null && info.model_name !== undefined && info.model_name !== '') { + message.model_name = info.model_name; + } // 只有在 status 不为空时更新 - if (info.status !== null && info.status !== undefined && info.status !== '') { - message.status = info.status; - } + if (info.status !== null && info.status !== undefined && info.status !== '') { + message.status = info.status; + } - if (info.meta !== null && info.meta !== undefined) { - message.meta = info.meta; + if (info.meta !== null && info.meta !== undefined) { + message.meta = info.meta; + } + scrollToBottom(); + } catch (error) { + console.error('Error updating message:', error); + message.status = 'error'; + message.text = '消息更新失败'; } } else { - console.error('Message not found'); + console.error('Message not found:', info.id); } - - scrollToBottom(); }; @@ -378,7 +389,7 @@ const groupRefs = (id) => { const simpleCall = (message) => { return new Promise((resolve, reject) => { - fetch('/api/chat/call', { + fetch('/api/chat/call_lite', { method: 'POST', body: JSON.stringify({ query: message, }), headers: { 'Content-Type': 'application/json' } @@ -457,10 +468,12 @@ const fetchChatResponse = (user_input, cur_res_id) => { updateMessage({ id: cur_res_id, text: data.response, + reasoning: data.reasoning_response, model_name: data.model_name, status: data.status, meta: data.meta, }); + // 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) @@ -648,6 +661,17 @@ watch( &:hover { background-color: var(--main-light-3); } + + .anticon { + margin-right: 8px; + font-size: 16px; + } + + .ant-switch { + &.ant-switch-checked { + background-color: var(--main-500); + } + } } } @@ -977,7 +1001,15 @@ watch( } } +.controls { + display: flex; + align-items: center; + gap: 8px; + .search-switch { + margin-right: 8px; + } +} + diff --git a/web/src/router/index.js b/web/src/router/index.js index 12e31833..35574dd3 100644 --- a/web/src/router/index.js +++ b/web/src/router/index.js @@ -11,7 +11,7 @@ const router = createRouter({ component: BlankLayout, children: [ { path: '', - name: 'home', + name: 'Home', component: () => import('../views/HomeView.vue'), meta: { keepAlive: true } } diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index d07e539f..132c882e 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -132,7 +132,7 @@

重写查询(修改后需重新检索)

- +