fix(embedding): 修复ollama model 不起作用的 bug,并修复model info 由于 getattr 不生效的 bug # 361
- 在 EmbedModelInfo 中添加 model_id 字段 - 修改 get_embedding_config 以优先使用 model_id 选择模型 - 重构 MilvusKB 的嵌入函数获取逻辑 - 添加调试日志并优化错误处理
This commit is contained in:
parent
001751dd46
commit
49253cbc04
@ -87,10 +87,12 @@ async def create_database(
|
||||
"""创建知识库"""
|
||||
logger.debug(
|
||||
f"Create database {database_name} with kb_type {kb_type}, "
|
||||
f"additional_params {additional_params}, llm_info {llm_info}"
|
||||
f"additional_params {additional_params}, llm_info {llm_info}, "
|
||||
f"embed_model_name {embed_model_name}"
|
||||
)
|
||||
try:
|
||||
additional_params = {**(additional_params or {})}
|
||||
additional_params["auto_generate_questions"] = False # 默认不生成问题
|
||||
|
||||
def normalize_reranker_config(kb: str, params: dict) -> None:
|
||||
reranker_cfg = params.get("reranker_config")
|
||||
@ -112,12 +114,12 @@ async def create_database(
|
||||
if not isinstance(reranker_cfg, Mapping):
|
||||
raise HTTPException(status_code=400, detail="reranker_config must be an object")
|
||||
|
||||
enabled = bool(reranker_cfg.get("enabled", False))
|
||||
reranker_enabled = bool(reranker_cfg.get("enabled", False))
|
||||
model = (reranker_cfg.get("model") or "").strip()
|
||||
recall_top_k = max(1, int(reranker_cfg.get("recall_top_k", 50)))
|
||||
final_top_k = max(1, int(reranker_cfg.get("final_top_k", 10)))
|
||||
|
||||
if enabled:
|
||||
if reranker_enabled:
|
||||
if not model:
|
||||
raise HTTPException(status_code=400, detail="reranker_config.model is required when enabled")
|
||||
if model not in config.reranker_names:
|
||||
@ -132,7 +134,7 @@ async def create_database(
|
||||
model = model if model in config.reranker_names else ""
|
||||
|
||||
params["reranker_config"] = {
|
||||
"enabled": enabled,
|
||||
"enabled": reranker_enabled,
|
||||
"model": model,
|
||||
"recall_top_k": recall_top_k,
|
||||
"final_top_k": final_top_k,
|
||||
|
||||
@ -29,7 +29,7 @@ class EmbedModelInfo(BaseModel):
|
||||
dimension: int = Field(..., description="向量维度")
|
||||
base_url: str = Field(..., description="API 基础 URL")
|
||||
api_key: str = Field(..., description="API Key 或环境变量名")
|
||||
|
||||
model_id: str | None = Field(None, description="可选的模型 ID")
|
||||
|
||||
class RerankerInfo(BaseModel):
|
||||
"""重排序模型配置"""
|
||||
@ -158,42 +158,49 @@ DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
|
||||
|
||||
DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
|
||||
"siliconflow/BAAI/bge-m3": EmbedModelInfo(
|
||||
model_id="siliconflow/BAAI/bge-m3",
|
||||
name="BAAI/bge-m3",
|
||||
dimension=1024,
|
||||
base_url="https://api.siliconflow.cn/v1/embeddings",
|
||||
api_key="SILICONFLOW_API_KEY",
|
||||
),
|
||||
"siliconflow/Pro/BAAI/bge-m3": EmbedModelInfo(
|
||||
model_id="siliconflow/Pro/BAAI/bge-m3",
|
||||
name="Pro/BAAI/bge-m3",
|
||||
dimension=1024,
|
||||
base_url="https://api.siliconflow.cn/v1/embeddings",
|
||||
api_key="SILICONFLOW_API_KEY",
|
||||
),
|
||||
"siliconflow/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
|
||||
model_id="siliconflow/Qwen/Qwen3-Embedding-0.6B",
|
||||
name="Qwen/Qwen3-Embedding-0.6B",
|
||||
dimension=1024,
|
||||
base_url="https://api.siliconflow.cn/v1/embeddings",
|
||||
api_key="SILICONFLOW_API_KEY",
|
||||
),
|
||||
"vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
|
||||
model_id="vllm/Qwen/Qwen3-Embedding-0.6B",
|
||||
name="Qwen3-Embedding-0.6B",
|
||||
dimension=1024,
|
||||
base_url="http://localhost:8000/v1/embeddings",
|
||||
api_key="no_api_key",
|
||||
),
|
||||
"ollama/nomic-embed-text": EmbedModelInfo(
|
||||
model_id="ollama/nomic-embed-text",
|
||||
name="nomic-embed-text",
|
||||
dimension=768,
|
||||
base_url="http://localhost:11434/api/embed",
|
||||
api_key="no_api_key",
|
||||
),
|
||||
"ollama/bge-m3": EmbedModelInfo(
|
||||
model_id="ollama/bge-m3",
|
||||
name="bge-m3",
|
||||
dimension=1024,
|
||||
base_url="http://localhost:11434/api/embed",
|
||||
api_key="no_api_key",
|
||||
),
|
||||
"dashscope/text-embedding-v4": EmbedModelInfo(
|
||||
model_id="dashscope/text-embedding-v4",
|
||||
name="text-embedding-v4",
|
||||
dimension=1024,
|
||||
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",
|
||||
|
||||
@ -7,6 +7,7 @@ from typing import Any
|
||||
|
||||
from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility
|
||||
|
||||
from src import config
|
||||
from src.knowledge.base import KnowledgeBase
|
||||
from src.knowledge.indexing import process_file_to_markdown
|
||||
from src.knowledge.utils.kb_utils import (
|
||||
@ -91,10 +92,14 @@ class MilvusKB(KnowledgeBase):
|
||||
"""创建 Milvus 集合"""
|
||||
logger.info(f"Creating Milvus collection for {db_id}")
|
||||
|
||||
if db_id not in self.databases_meta:
|
||||
if not (metadata := self.databases_meta.get(db_id)):
|
||||
raise ValueError(f"Database {db_id} not found")
|
||||
|
||||
embed_info = self.databases_meta[db_id].get("embed_info", {})
|
||||
# embed_info = metadata.get("embed_info", {})
|
||||
if not (embed_info := metadata.get("embed_info")):
|
||||
logger.error(f"Embedding info not found for database {db_id}, using default model")
|
||||
embed_info = config.embed_model_names[config.embed_model]
|
||||
|
||||
collection_name = db_id
|
||||
|
||||
try:
|
||||
@ -117,8 +122,8 @@ class MilvusKB(KnowledgeBase):
|
||||
|
||||
except Exception:
|
||||
# 创建新集合
|
||||
embedding_dim = getattr(embed_info, "dimension", 1024) if embed_info else 1024
|
||||
model_name = getattr(embed_info, "name", "default") if embed_info else "default"
|
||||
embedding_dim = embed_info.get("dimension", 1024)
|
||||
model_name = embed_info.get("name", "default")
|
||||
|
||||
# 定义集合Schema
|
||||
fields = [
|
||||
@ -142,7 +147,7 @@ class MilvusKB(KnowledgeBase):
|
||||
index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}}
|
||||
collection.create_index("embedding", index_params)
|
||||
|
||||
logger.info(f"Created new Milvus collection: {collection_name}")
|
||||
logger.info(f"Created new Milvus collection: {collection_name}: {model_name=}, {embedding_dim=}")
|
||||
|
||||
return collection
|
||||
|
||||
@ -154,25 +159,29 @@ class MilvusKB(KnowledgeBase):
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load collection into memory: {e}")
|
||||
|
||||
def _get_async_embedding_function(self, embed_info: dict):
|
||||
def _get_async_embedding(self, embed_info: dict):
|
||||
"""获取 embedding 函数"""
|
||||
# 检查是否有 model_id 字段,优先使用 select_embedding_model
|
||||
if embed_info and "model_id" in embed_info:
|
||||
from src.models.embed import select_embedding_model
|
||||
return select_embedding_model(embed_info["model_id"])
|
||||
|
||||
# 使用原有的逻辑(兼容模式))
|
||||
config_dict = get_embedding_config(embed_info)
|
||||
embedding_model = OtherEmbedding(
|
||||
return OtherEmbedding(
|
||||
model=config_dict.get("model"),
|
||||
base_url=config_dict.get("base_url"),
|
||||
api_key=config_dict.get("api_key"),
|
||||
)
|
||||
|
||||
def _get_async_embedding_function(self, embed_info: dict):
|
||||
"""获取 embedding 函数"""
|
||||
embedding_model = self._get_async_embedding(embed_info)
|
||||
return partial(embedding_model.abatch_encode, batch_size=40)
|
||||
|
||||
def _get_embedding_function(self, embed_info: dict):
|
||||
"""获取 embedding 函数"""
|
||||
config_dict = get_embedding_config(embed_info)
|
||||
embedding_model = OtherEmbedding(
|
||||
model=config_dict.get("model"),
|
||||
base_url=config_dict.get("base_url"),
|
||||
api_key=config_dict.get("api_key"),
|
||||
)
|
||||
embedding_model = self._get_async_embedding(embed_info)
|
||||
|
||||
return partial(embedding_model.batch_encode, batch_size=40)
|
||||
|
||||
|
||||
@ -246,16 +246,12 @@ class KnowledgeBaseManager:
|
||||
db_id = db_info["db_id"]
|
||||
|
||||
async with self._metadata_lock:
|
||||
# 准备 additional_params,包含 auto_generate_questions
|
||||
saved_params = kwargs.copy()
|
||||
saved_params["auto_generate_questions"] = False
|
||||
|
||||
self.global_databases_meta[db_id] = {
|
||||
"name": database_name,
|
||||
"description": description,
|
||||
"kb_type": kb_type,
|
||||
"created_at": utc_isoformat(),
|
||||
"additional_params": saved_params,
|
||||
"additional_params": kwargs.copy(),
|
||||
}
|
||||
self._save_global_metadata()
|
||||
|
||||
|
||||
@ -247,15 +247,23 @@ def get_embedding_config(embed_info: dict) -> dict:
|
||||
|
||||
try:
|
||||
if embed_info:
|
||||
# 处理 embed_info 可能是字典或 EmbedModelInfo 对象的情况
|
||||
if hasattr(embed_info, "name"):
|
||||
# 优先检查是否有 model_id 字段
|
||||
if "model_id" in embed_info:
|
||||
from src.models.embed import select_embedding_model
|
||||
|
||||
model = select_embedding_model(embed_info["model_id"])
|
||||
config_dict["model"] = model.model
|
||||
config_dict["api_key"] = model.api_key
|
||||
config_dict["base_url"] = model.base_url
|
||||
config_dict["dimension"] = getattr(model, "dimension", 1024)
|
||||
elif hasattr(embed_info, "name"):
|
||||
# EmbedModelInfo 对象
|
||||
config_dict["model"] = embed_info.name
|
||||
config_dict["api_key"] = os.getenv(embed_info.api_key) or embed_info.api_key
|
||||
config_dict["base_url"] = embed_info.base_url
|
||||
config_dict["dimension"] = embed_info.dimension
|
||||
else:
|
||||
# 字典形式
|
||||
# 字典形式(保持向后兼容)
|
||||
config_dict["model"] = embed_info["name"]
|
||||
config_dict["api_key"] = os.getenv(embed_info["api_key"]) or embed_info["api_key"]
|
||||
config_dict["base_url"] = embed_info["base_url"]
|
||||
|
||||
@ -11,7 +11,7 @@ from src.utils import get_docker_safe_url, hashstr, logger
|
||||
|
||||
|
||||
class BaseEmbeddingModel(ABC):
|
||||
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None):
|
||||
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None, model_id=None):
|
||||
"""
|
||||
Args:
|
||||
model: 模型名称,冗余设计,同name
|
||||
@ -140,6 +140,7 @@ class OllamaEmbedding(BaseEmbeddingModel):
|
||||
payload = {"model": self.model, "input": message}
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
print(f"\n\n\nOllama Embedding request: {payload}\n\n\n")
|
||||
response = await client.post(self.base_url, json=payload, timeout=60)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user