Merge branch 'main' of https://github.com/xerrors/Yuxi-Know
This commit is contained in:
commit
9e46215379
43
README.md
43
README.md
@ -1,4 +1,4 @@
|
||||
<h1 align="center">语析 - 基于大模型的知识库与知识图谱问答平台</h1>
|
||||
<h1 align="center">语析 - 基于大模型的知识库与知识图谱问答系统</h1>
|
||||
<div align="center">
|
||||
|
||||

|
||||
@ -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)
|
||||
|
||||
@ -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}
|
||||
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
@ -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()}")
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
@ -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"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user