diff --git a/README.md b/README.md
index 030b8486..c64ea233 100644
--- a/README.md
+++ b/README.md
@@ -1,4 +1,4 @@
-
语析 - 基于大模型的知识库与知识图谱问答平台
+语析 - 基于大模型的知识库与知识图谱问答系统

@@ -23,7 +23,6 @@

-
## 📋 更新日志
- **2025.02.24** - 新增网页检索以及内容展示,需配置 `TAVILY_API_KEY`,感谢 [littlewwwhite](https://github.com/littlewwwhite)
@@ -33,11 +32,9 @@

-
-| PC 网页 | 小屏设备 |
-|:-----------|:-----------|
-| | |
-
+| PC 网页 | 小屏设备 |
+| :-------------------------------------------------------------------------------------- | :-------------------------------------------------------------------------------------- |
+|  |  |
### 环境配置
@@ -149,7 +146,33 @@ ark:
对于**向量模型**和**重排序模型**,选择 `local` 前缀的模型会自动下载。如遇下载问题,请参考 [HF-Mirror](https://hf-mirror.com/) 配置。
-要使用已下载的本地模型,可在网页设置中映射,或修改 `saves/config/base.yaml`。记得在 docker-compose 中映射相应的 volumes。
+要使用已下载的本地模型,可在 models.yaml 或者网页设置中映射。
+
+
+
+
+**添加向量模型**
+
+```yaml
+# src/static/models.yaml
+ # 添加本地向量模型(所有 FlagEmbedding 支持的模型)
+ local/BAAI/bge-m3:
+ name: BAAI/bge-m3
+ dimension: 1024
+ # local_path: /models/BAAI/bge-m3,也可以在这里配置
+
+ # 添加 OpenAI 兼容的向量模型
+ siliconflow/BAAI/bge-m3:
+ name: BAAI/bge-m3
+ dimension: 1024
+ url: https://api.siliconflow.cn/v1/embeddings
+ api_key: SILICONFLOW_API_KEY
+
+ # 添加 Ollama 模型
+ ollama/nomic-embed-text:
+ name: nomic-embed-text
+ dimension: 768
+```
## 📚 知识库支持
@@ -204,3 +227,7 @@ docker pull m.daocloud.io/docker.io/library/neo4j:latest
# 然后重命名镜像
docker tag m.daocloud.io/docker.io/library/neo4j:latest neo4j:latest
```
+
+## Star History
+
+[](https://star-history.com/#xerrors/Yuxi-Know)
diff --git a/src/config/__init__.py b/src/config/__init__.py
index 2de995f0..9ebdb38d 100644
--- a/src/config/__init__.py
+++ b/src/config/__init__.py
@@ -82,6 +82,8 @@ class Config(SimpleConfig):
"_config_items",
"model_names",
"model_provider_status",
+ "embed_model_names",
+ "reranker_names",
]
return {k: v for k, v in self.items() if k not in blocklist}
diff --git a/src/models/chat_model.py b/src/models/chat_model.py
index d146b3bf..d85826a0 100644
--- a/src/models/chat_model.py
+++ b/src/models/chat_model.py
@@ -1,6 +1,6 @@
import os
from openai import OpenAI
-from src.utils import logger
+from src.utils import logger, get_docker_safe_url
class OpenAIBase():
def __init__(self, api_key, base_url, model_name):
@@ -44,14 +44,6 @@ class OpenModel(OpenAIBase):
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
-def get_docker_safe_url(base_url):
- if os.getenv("RUNNING_IN_DOCKER") == "true":
- # 替换所有可能的本地地址形式
- base_url = base_url.replace("http://localhost", "http://host.docker.internal")
- base_url = base_url.replace("http://127.0.0.1", "http://host.docker.internal")
- logger.info(f"Running in docker, using {base_url} as base url")
- return base_url
-
class CustomModel(OpenAIBase):
def __init__(self, model_info):
diff --git a/src/models/embedding.py b/src/models/embedding.py
index e7d7b6b4..d38129c1 100644
--- a/src/models/embedding.py
+++ b/src/models/embedding.py
@@ -5,14 +5,21 @@ from FlagEmbedding import FlagModel
from zhipuai import ZhipuAI
from src.config import EMBED_MODEL_INFO
-from src.utils import hashstr, logger
+from src.utils import hashstr, logger, get_docker_safe_url
-class RemoteEmbeddingModel:
+class BaseEmbeddingModel:
embed_state = {}
+ EMBED_MODEL_INFO = EMBED_MODEL_INFO
+
+ def encode(self, message):
+ return self.predict(message)
+
+ def encode_queries(self, queries):
+ return self.predict(queries)
def batch_encode(self, messages, batch_size=20):
- logger.info(f"Batch encoding {len(messages)} messages")
+ logger.info(f"[Embedding: {self.model}] Batch encoding {len(messages)} messages")
data = []
if len(messages) > batch_size:
@@ -35,45 +42,23 @@ class RemoteEmbeddingModel:
return data
-class LocalEmbeddingModel(FlagModel, RemoteEmbeddingModel):
+class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel):
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,
+ self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
+ self.model = self.model or info["name"]
+
+ logger.info(f"Loading embedding model {info['name']} from {self.model}")
+
+ super().__init__(self.model,
query_instruction_for_retrieval=info.get("query_instruction", None),
use_fp16=False, **kwargs)
logger.info(f"Embedding model {info['name']} loaded")
- def batch_encode(self, messages, batch_size=20):
- logger.info(f"Batch encoding {len(messages)} messages")
- data = []
-
- if len(messages) > batch_size:
- task_id = hashstr(messages)
- self.embed_state[task_id] = {
- 'status': 'in-progress',
- 'total': len(messages),
- 'progress': 0
- }
-
- 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)
-
- if len(messages) > batch_size:
- self.embed_state[task_id]['progress'] = len(messages)
- self.embed_state[task_id]['status'] = 'completed'
-
- return data
-
-
-class ZhipuEmbedding(RemoteEmbeddingModel):
+class ZhipuEmbedding(BaseEmbeddingModel):
def __init__(self, config) -> None:
self.config = config
@@ -88,37 +73,49 @@ class ZhipuEmbedding(RemoteEmbeddingModel):
data = [a.embedding for a in response.data]
return data
- def encode(self, message):
- return self.predict(message)
- def encode_queries(self, queries):
- return self.predict(queries)
+class OllamaEmbedding(BaseEmbeddingModel):
+ def __init__(self, config) -> None:
+ self.info = EMBED_MODEL_INFO[config.embed_model]
+ self.model = self.info["name"]
+ self.url = self.info.get("url", "http://localhost:11434/api/embed")
+ self.url = get_docker_safe_url(self.url)
+
+ def predict(self, message: list[str] | str):
+ if isinstance(message, str):
+ message = [message]
+
+ payload = {
+ "model": self.model,
+ "input": message,
+ }
+ response = requests.request("POST", self.url, json=payload)
+ response = json.loads(response.text)
+ assert response.get("embeddings"), f"Ollama Embedding failed: {response}"
+ return response["embeddings"]
-class SiliconFlowEmbedding(RemoteEmbeddingModel):
+class OtherEmbedding(BaseEmbeddingModel):
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.info = EMBED_MODEL_INFO[config.embed_model]
+ self.model = self.info["name"]
+ self.api_key = os.getenv(self.info["api_key"], None)
+ self.url = get_docker_safe_url(self.info["url"])
+ assert self.url and self.model, f"URL and model are required. Cur embed model: {config.embed_model}"
self.headers = {
- "Authorization": f"Bearer {api_key}",
+ "Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
}
- def encode(self, message):
+ def predict(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}"
+ assert response["data"], f"Other 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,
@@ -135,11 +132,14 @@ def get_embedding_model(config):
if provider == "local":
model = LocalEmbeddingModel(config)
- if provider == "zhipu":
+ elif provider == "zhipu":
model = ZhipuEmbedding(config)
- if provider == "siliconflow":
- model = SiliconFlowEmbedding(config)
+ elif provider == "ollama":
+ model = OllamaEmbedding(config)
+
+ else:
+ model = OtherEmbedding(config)
return model
diff --git a/src/models/ollama_embedding.py b/src/models/ollama_embedding.py
deleted file mode 100644
index 1aefc16c..00000000
--- a/src/models/ollama_embedding.py
+++ /dev/null
@@ -1,179 +0,0 @@
-import os
-import requests
-import numpy as np
-from typing import List, Union, Dict
-
-from src.models.embedding import RemoteEmbeddingModel
-from src.utils.logging_config import logger
-
-class OllamaEmbedding(RemoteEmbeddingModel):
- """
- 使用 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/models/rerank_model.py b/src/models/rerank_model.py
index 5f02a75e..d1cf1136 100644
--- a/src/models/rerank_model.py
+++ b/src/models/rerank_model.py
@@ -11,7 +11,8 @@ from src.utils.logging_config import logger
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"))
+ model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("local_path"))
+ model_name_or_path = model_name_or_path or model_info["name"]
logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}")
super().__init__(model_name_or_path, use_fp16=True, **kwargs)
@@ -64,4 +65,6 @@ def get_reranker(config):
return LocalReranker(config)
elif provider == "siliconflow":
return SilconFlowReranker(config)
+ else:
+ raise ValueError(f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}")
diff --git a/src/static/models.yaml b/src/static/models.yaml
index 329f9f54..a7a813c3 100644
--- a/src/static/models.yaml
+++ b/src/static/models.yaml
@@ -114,21 +114,36 @@ MODEL_NAMES:
EMBED_MODEL_INFO:
local/BAAI/bge-m3:
name: BAAI/bge-m3
- default_path: BAAI/bge-m3
dimension: 1024
+ # local_path: /models/BAAI/bge-m3,也可以在这里配置
+
zhipu/zhipu-embedding-2:
name: embedding-2
dimension: 1024
+
zhipu/zhipu-embedding-3:
name: embedding-3
dimension: 2048
+
siliconflow/BAAI/bge-m3:
name: BAAI/bge-m3
dimension: 1024
+ url: https://api.siliconflow.cn/v1/embeddings
+ api_key: SILICONFLOW_API_KEY
+
+ ollama/nomic-embed-text:
+ name: nomic-embed-text
+ dimension: 768
+
+ ollama/bge-m3:
+ name: bge-m3
+ dimension: 1024
RERANKER_LIST:
+
local/BAAI/bge-reranker-v2-m3:
name: BAAI/bge-reranker-v2-m3
- default_path: BAAI/bge-reranker-v2-m3
+ # local_path: /models/BAAI/bge-m3,也可以在这里配置
+
siliconflow/BAAI/bge-reranker-v2-m3:
name: BAAI/bge-reranker-v2-m3
diff --git a/src/utils/__init__.py b/src/utils/__init__.py
index 37ff7c4a..39882d03 100644
--- a/src/utils/__init__.py
+++ b/src/utils/__init__.py
@@ -1,5 +1,6 @@
import time
import random
+import os
from src.utils.logging_config import logger
def is_text_pdf(pdf_path):
@@ -19,4 +20,13 @@ def hashstr(input_string, length=8, with_salt=False):
input_string += str(time.time() + random.random())
hash = hashlib.md5(str(input_string).encode()).hexdigest()
- return hash[:length]
\ No newline at end of file
+ return hash[:length]
+
+
+def get_docker_safe_url(base_url):
+ if os.getenv("RUNNING_IN_DOCKER") == "true":
+ # 替换所有可能的本地地址形式
+ base_url = base_url.replace("http://localhost", "http://host.docker.internal")
+ base_url = base_url.replace("http://127.0.0.1", "http://host.docker.internal")
+ logger.info(f"Running in docker, using {base_url} as base url")
+ return base_url
\ No newline at end of file
diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue
index 6d22b3f0..1e307d59 100644
--- a/web/src/views/SettingView.vue
+++ b/web/src/views/SettingView.vue
@@ -186,7 +186,7 @@
本地模型配置
-
如果是 Docker 启动,务必确保在环境变量中设置了 MODEL_DIR,或者设置了 volumes 映射
+
如果是 Docker 启动,务必确保在环境变量中设置了 MODEL_DIR(建议是绝对路径),并设置了 volumes 映射。