This commit is contained in:
Wenjie Zhang 2025-03-05 16:31:42 +08:00
commit 9e46215379
9 changed files with 123 additions and 253 deletions

View File

@ -1,4 +1,4 @@
<h1 align="center">语析 - 基于大模型的知识库与知识图谱问答平台</h1>
<h1 align="center">语析 - 基于大模型的知识库与知识图谱问答系统</h1>
<div align="center">
![](https://img.shields.io/badge/Docker-2496ED?style=flat&logo=docker&logoColor=ffffff)
@ -23,7 +23,6 @@
![系统界面预览](https://github.com/user-attachments/assets/75010511-4ac5-4924-8268-fea9a589839c)
## 📋 更新日志
- **2025.02.24** - 新增网页检索以及内容展示,需配置 `TAVILY_API_KEY`,感谢 [littlewwwhite](https://github.com/littlewwwhite)
@ -33,11 +32,9 @@
![功能展示](https://github.com/user-attachments/assets/8416a933-cc43-45d0-bf06-00df0ba6c4fb)
| PC 网页 | 小屏设备 |
|:-----------|:-----------|
| ![image](https://github.com/user-attachments/assets/5f3d7e69-baa8-4c59-90fc-391343e59af6)| ![image](https://github.com/user-attachments/assets/51efabce-a097-47fd-9fca-d3b0943af86a)|
| PC 网页 | 小屏设备 |
| :-------------------------------------------------------------------------------------- | :-------------------------------------------------------------------------------------- |
| ![image](https://github.com/user-attachments/assets/5f3d7e69-baa8-4c59-90fc-391343e59af6) | ![image](https://github.com/user-attachments/assets/51efabce-a097-47fd-9fca-d3b0943af86a) |
### 环境配置
@ -149,7 +146,33 @@ ark:
对于**向量模型**和**重排序模型**,选择 `local` 前缀的模型会自动下载。如遇下载问题,请参考 [HF-Mirror](https://hf-mirror.com/) 配置。
要使用已下载的本地模型,可在网页设置中映射,或修改 `saves/config/base.yaml`。记得在 docker-compose 中映射相应的 volumes。
要使用已下载的本地模型,可在 models.yaml 或者网页设置中映射。
![image](https://github.com/user-attachments/assets/ab62ea17-c7d0-4f94-84af-c4bab26865ad)
**添加向量模型**
```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
[![Star History Chart](https://api.star-history.com/svg?repos=xerrors/Yuxi-Know)](https://star-history.com/#xerrors/Yuxi-Know)

View File

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

View File

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

View File

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

View File

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

View File

@ -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()}")

View File

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

View File

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

View File

@ -186,7 +186,7 @@
</div>
<div class="setting" v-if="state.windowWidth <= 520 || state.section ==='path'">
<h3>本地模型配置</h3>
<p>如果是 Docker 启动务必确保在环境变量中设置了 MODEL_DIR或者设置了 volumes 映射</p>
<p>如果是 Docker 启动务必确保在环境变量中设置了 MODEL_DIR建议是绝对路径并设置了 volumes 映射</p>
<TableConfigComponent
:config="configStore.config?.model_local_paths"
@update:config="handleModelLocalPathsUpdate"