fix(embedding): 修复ollama model 不起作用的 bug,并修复model info 由于 getattr 不生效的 bug # 361

- 在 EmbedModelInfo 中添加 model_id 字段
- 修改 get_embedding_config 以优先使用 model_id 选择模型
- 重构 MilvusKB 的嵌入函数获取逻辑
- 添加调试日志并优化错误处理
This commit is contained in:
Wenjie Zhang 2025-12-03 18:28:31 +08:00
parent 001751dd46
commit 49253cbc04
6 changed files with 50 additions and 27 deletions

View File

@ -87,10 +87,12 @@ async def create_database(
"""创建知识库""" """创建知识库"""
logger.debug( logger.debug(
f"Create database {database_name} with kb_type {kb_type}, " 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: try:
additional_params = {**(additional_params or {})} additional_params = {**(additional_params or {})}
additional_params["auto_generate_questions"] = False # 默认不生成问题
def normalize_reranker_config(kb: str, params: dict) -> None: def normalize_reranker_config(kb: str, params: dict) -> None:
reranker_cfg = params.get("reranker_config") reranker_cfg = params.get("reranker_config")
@ -112,12 +114,12 @@ async def create_database(
if not isinstance(reranker_cfg, Mapping): if not isinstance(reranker_cfg, Mapping):
raise HTTPException(status_code=400, detail="reranker_config must be an object") 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() model = (reranker_cfg.get("model") or "").strip()
recall_top_k = max(1, int(reranker_cfg.get("recall_top_k", 50))) 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))) final_top_k = max(1, int(reranker_cfg.get("final_top_k", 10)))
if enabled: if reranker_enabled:
if not model: if not model:
raise HTTPException(status_code=400, detail="reranker_config.model is required when enabled") raise HTTPException(status_code=400, detail="reranker_config.model is required when enabled")
if model not in config.reranker_names: if model not in config.reranker_names:
@ -132,7 +134,7 @@ async def create_database(
model = model if model in config.reranker_names else "" model = model if model in config.reranker_names else ""
params["reranker_config"] = { params["reranker_config"] = {
"enabled": enabled, "enabled": reranker_enabled,
"model": model, "model": model,
"recall_top_k": recall_top_k, "recall_top_k": recall_top_k,
"final_top_k": final_top_k, "final_top_k": final_top_k,

View File

@ -29,7 +29,7 @@ class EmbedModelInfo(BaseModel):
dimension: int = Field(..., description="向量维度") dimension: int = Field(..., description="向量维度")
base_url: str = Field(..., description="API 基础 URL") base_url: str = Field(..., description="API 基础 URL")
api_key: str = Field(..., description="API Key 或环境变量名") api_key: str = Field(..., description="API Key 或环境变量名")
model_id: str | None = Field(None, description="可选的模型 ID")
class RerankerInfo(BaseModel): class RerankerInfo(BaseModel):
"""重排序模型配置""" """重排序模型配置"""
@ -158,42 +158,49 @@ DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = { DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
"siliconflow/BAAI/bge-m3": EmbedModelInfo( "siliconflow/BAAI/bge-m3": EmbedModelInfo(
model_id="siliconflow/BAAI/bge-m3",
name="BAAI/bge-m3", name="BAAI/bge-m3",
dimension=1024, dimension=1024,
base_url="https://api.siliconflow.cn/v1/embeddings", base_url="https://api.siliconflow.cn/v1/embeddings",
api_key="SILICONFLOW_API_KEY", api_key="SILICONFLOW_API_KEY",
), ),
"siliconflow/Pro/BAAI/bge-m3": EmbedModelInfo( "siliconflow/Pro/BAAI/bge-m3": EmbedModelInfo(
model_id="siliconflow/Pro/BAAI/bge-m3",
name="Pro/BAAI/bge-m3", name="Pro/BAAI/bge-m3",
dimension=1024, dimension=1024,
base_url="https://api.siliconflow.cn/v1/embeddings", base_url="https://api.siliconflow.cn/v1/embeddings",
api_key="SILICONFLOW_API_KEY", api_key="SILICONFLOW_API_KEY",
), ),
"siliconflow/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo( "siliconflow/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
model_id="siliconflow/Qwen/Qwen3-Embedding-0.6B",
name="Qwen/Qwen3-Embedding-0.6B", name="Qwen/Qwen3-Embedding-0.6B",
dimension=1024, dimension=1024,
base_url="https://api.siliconflow.cn/v1/embeddings", base_url="https://api.siliconflow.cn/v1/embeddings",
api_key="SILICONFLOW_API_KEY", api_key="SILICONFLOW_API_KEY",
), ),
"vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo( "vllm/Qwen/Qwen3-Embedding-0.6B": EmbedModelInfo(
model_id="vllm/Qwen/Qwen3-Embedding-0.6B",
name="Qwen3-Embedding-0.6B", name="Qwen3-Embedding-0.6B",
dimension=1024, dimension=1024,
base_url="http://localhost:8000/v1/embeddings", base_url="http://localhost:8000/v1/embeddings",
api_key="no_api_key", api_key="no_api_key",
), ),
"ollama/nomic-embed-text": EmbedModelInfo( "ollama/nomic-embed-text": EmbedModelInfo(
model_id="ollama/nomic-embed-text",
name="nomic-embed-text", name="nomic-embed-text",
dimension=768, dimension=768,
base_url="http://localhost:11434/api/embed", base_url="http://localhost:11434/api/embed",
api_key="no_api_key", api_key="no_api_key",
), ),
"ollama/bge-m3": EmbedModelInfo( "ollama/bge-m3": EmbedModelInfo(
model_id="ollama/bge-m3",
name="bge-m3", name="bge-m3",
dimension=1024, dimension=1024,
base_url="http://localhost:11434/api/embed", base_url="http://localhost:11434/api/embed",
api_key="no_api_key", api_key="no_api_key",
), ),
"dashscope/text-embedding-v4": EmbedModelInfo( "dashscope/text-embedding-v4": EmbedModelInfo(
model_id="dashscope/text-embedding-v4",
name="text-embedding-v4", name="text-embedding-v4",
dimension=1024, dimension=1024,
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings", base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",

View File

@ -7,6 +7,7 @@ from typing import Any
from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility
from src import config
from src.knowledge.base import KnowledgeBase from src.knowledge.base import KnowledgeBase
from src.knowledge.indexing import process_file_to_markdown from src.knowledge.indexing import process_file_to_markdown
from src.knowledge.utils.kb_utils import ( from src.knowledge.utils.kb_utils import (
@ -91,10 +92,14 @@ class MilvusKB(KnowledgeBase):
"""创建 Milvus 集合""" """创建 Milvus 集合"""
logger.info(f"Creating Milvus collection for {db_id}") 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") 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 collection_name = db_id
try: try:
@ -117,8 +122,8 @@ class MilvusKB(KnowledgeBase):
except Exception: except Exception:
# 创建新集合 # 创建新集合
embedding_dim = getattr(embed_info, "dimension", 1024) if embed_info else 1024 embedding_dim = embed_info.get("dimension", 1024)
model_name = getattr(embed_info, "name", "default") if embed_info else "default" model_name = embed_info.get("name", "default")
# 定义集合Schema # 定义集合Schema
fields = [ fields = [
@ -142,7 +147,7 @@ class MilvusKB(KnowledgeBase):
index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}} index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}}
collection.create_index("embedding", index_params) 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 return collection
@ -154,25 +159,29 @@ class MilvusKB(KnowledgeBase):
except Exception as e: except Exception as e:
logger.warning(f"Failed to load collection into memory: {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 函数""" """获取 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) config_dict = get_embedding_config(embed_info)
embedding_model = OtherEmbedding( return OtherEmbedding(
model=config_dict.get("model"), model=config_dict.get("model"),
base_url=config_dict.get("base_url"), base_url=config_dict.get("base_url"),
api_key=config_dict.get("api_key"), 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) return partial(embedding_model.abatch_encode, batch_size=40)
def _get_embedding_function(self, embed_info: dict): def _get_embedding_function(self, embed_info: dict):
"""获取 embedding 函数""" """获取 embedding 函数"""
config_dict = get_embedding_config(embed_info) embedding_model = self._get_async_embedding(embed_info)
embedding_model = OtherEmbedding(
model=config_dict.get("model"),
base_url=config_dict.get("base_url"),
api_key=config_dict.get("api_key"),
)
return partial(embedding_model.batch_encode, batch_size=40) return partial(embedding_model.batch_encode, batch_size=40)

View File

@ -246,16 +246,12 @@ class KnowledgeBaseManager:
db_id = db_info["db_id"] db_id = db_info["db_id"]
async with self._metadata_lock: 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] = { self.global_databases_meta[db_id] = {
"name": database_name, "name": database_name,
"description": description, "description": description,
"kb_type": kb_type, "kb_type": kb_type,
"created_at": utc_isoformat(), "created_at": utc_isoformat(),
"additional_params": saved_params, "additional_params": kwargs.copy(),
} }
self._save_global_metadata() self._save_global_metadata()

View File

@ -247,15 +247,23 @@ def get_embedding_config(embed_info: dict) -> dict:
try: try:
if embed_info: if embed_info:
# 处理 embed_info 可能是字典或 EmbedModelInfo 对象的情况 # 优先检查是否有 model_id 字段
if hasattr(embed_info, "name"): 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 对象 # EmbedModelInfo 对象
config_dict["model"] = embed_info.name config_dict["model"] = embed_info.name
config_dict["api_key"] = os.getenv(embed_info.api_key) or embed_info.api_key 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["base_url"] = embed_info.base_url
config_dict["dimension"] = embed_info.dimension config_dict["dimension"] = embed_info.dimension
else: else:
# 字典形式 # 字典形式(保持向后兼容)
config_dict["model"] = embed_info["name"] config_dict["model"] = embed_info["name"]
config_dict["api_key"] = os.getenv(embed_info["api_key"]) or embed_info["api_key"] 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["base_url"] = embed_info["base_url"]

View File

@ -11,7 +11,7 @@ from src.utils import get_docker_safe_url, hashstr, logger
class BaseEmbeddingModel(ABC): 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: Args:
model: 模型名称冗余设计同name model: 模型名称冗余设计同name
@ -140,6 +140,7 @@ class OllamaEmbedding(BaseEmbeddingModel):
payload = {"model": self.model, "input": message} payload = {"model": self.model, "input": message}
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
try: try:
print(f"\n\n\nOllama Embedding request: {payload}\n\n\n")
response = await client.post(self.base_url, json=payload, timeout=60) response = await client.post(self.base_url, json=payload, timeout=60)
response.raise_for_status() response.raise_for_status()
result = response.json() result = response.json()