From 8198f2f39024034d72ba70a4b2952c34c7634809 Mon Sep 17 00:00:00 2001 From: lintsinghua Date: Sat, 24 Jan 2026 14:42:31 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=B5=8C=E5=85=A5=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E5=90=91=E9=87=8F=E7=BB=B4=E5=BA=A6=E5=8F=AF=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=EF=BC=8C=E8=A7=A3=E5=86=B3=20Ollama=20=E4=B8=8D=E5=90=8C?= =?UTF-8?q?=E5=8F=82=E6=95=B0=E8=A7=84=E6=A8=A1=E6=A8=A1=E5=9E=8B=E7=BB=B4?= =?UTF-8?q?=E5=BA=A6=E4=B8=8D=E5=8C=B9=E9=85=8D=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #123 问题:qwen3-embedding:8b 实际维度 4096,但代码硬编码为 1024(只适用于 0.6b 版本),导致 RAG 系统初始化失败 修改内容: - OllamaEmbedding: 构造函数添加 dimension 参数,用户配置优先 - EmbeddingService: 支持传递自定义维度到各提供商 - embedding_config.py: - 补充 _get_model_dimensions 缺失的模型映射 - get_current_config 优先使用用户配置的维度 - test_embedding 支持传递自定义维度 - 前端 EmbeddingConfig: 添加"自定义向量维度"输入框 使用方式: 在"系统配置-嵌入模型"页面,输入自定义向量维度(如 4096)即可覆盖默认值 Co-Authored-By: Claude Opus 4.5 --- .../app/api/v1/endpoints/embedding_config.py | 30 ++++++++----- backend/app/services/rag/embeddings.py | 43 ++++++++++++------- .../src/components/agent/EmbeddingConfig.tsx | 25 +++++++++++ 3 files changed, 71 insertions(+), 27 deletions(-) diff --git a/backend/app/api/v1/endpoints/embedding_config.py b/backend/app/api/v1/endpoints/embedding_config.py index 3760bd9..4a1e68d 100644 --- a/backend/app/api/v1/endpoints/embedding_config.py +++ b/backend/app/api/v1/endpoints/embedding_config.py @@ -61,6 +61,7 @@ class TestEmbeddingRequest(BaseModel): model: str api_key: Optional[str] = None base_url: Optional[str] = None + dimension: Optional[int] = None # 自定义维度(Ollama等场景) test_text: str = "这是一段测试文本,用于验证嵌入模型是否正常工作。" @@ -275,8 +276,8 @@ async def get_current_config( """ config = await get_embedding_config_from_db(db, current_user.id) - # 获取维度 - dimensions = _get_model_dimensions(config.provider, config.model) + # 获取维度:优先使用用户配置的维度,否则使用默认值 + dimensions = config.dimensions if config.dimensions else _get_model_dimensions(config.provider, config.model) return EmbeddingConfigResponse( provider=config.provider, @@ -326,18 +327,19 @@ async def test_embedding( """ FIXED_DURATION = 3.0 # 固定响应时间,防止SSRF时间侧信道攻击 start_time = time.time() - + try: from app.services.rag.embeddings import EmbeddingService - + service = EmbeddingService( provider=request.provider, model=request.model, api_key=request.api_key, base_url=request.base_url, + dimension=request.dimension, cache_enabled=False, ) - + embedding = await service.embed(request.test_text) elapsed = time.time() - start_time @@ -393,35 +395,41 @@ def _get_model_dimensions(provider: str, model: str) -> int: "text-embedding-3-small": 1536, "text-embedding-3-large": 3072, "text-embedding-ada-002": 1536, - + # Ollama "nomic-embed-text": 768, "mxbai-embed-large": 1024, "all-minilm": 384, "snowflake-arctic-embed": 1024, - + "bge-m3": 1024, + "qwen3-embedding": 1024, # 默认值,8b版本为4096 + # Cohere "embed-english-v3.0": 1024, "embed-multilingual-v3.0": 1024, "embed-english-light-v3.0": 384, "embed-multilingual-light-v3.0": 384, - + "embed-v4.0": 1024, + # HuggingFace "sentence-transformers/all-MiniLM-L6-v2": 384, "sentence-transformers/all-mpnet-base-v2": 768, "BAAI/bge-large-zh-v1.5": 1024, "BAAI/bge-m3": 1024, - + "BAAI/bge-small-en-v1.5": 384, + "BAAI/bge-base-en-v1.5": 768, + # Jina "jina-embeddings-v2-base-code": 768, "jina-embeddings-v2-base-en": 768, "jina-embeddings-v2-base-zh": 768, - + "jina-embeddings-v2-small-en": 512, + # Qwen (DashScope) "text-embedding-v4": 1024, # 支持维度: 2048, 1536, 1024(默认), 768, 512, 256, 128, 64 "text-embedding-v3": 1024, # 支持维度: 1024(默认), 768, 512, 256, 128, 64 "text-embedding-v2": 1536, # 支持维度: 1536 } - + return dimensions_map.get(model, 768) diff --git a/backend/app/services/rag/embeddings.py b/backend/app/services/rag/embeddings.py index 9682f9c..8cb5d28 100644 --- a/backend/app/services/rag/embeddings.py +++ b/backend/app/services/rag/embeddings.py @@ -182,29 +182,35 @@ class AzureOpenAIEmbedding(EmbeddingProvider): class OllamaEmbedding(EmbeddingProvider): """ Ollama 本地嵌入服务 - + 使用新的 /api/embed 端点 (2024年起): - 支持批量嵌入 - 使用 'input' 参数(支持字符串或字符串数组) """ - + + # 默认维度映射(基础模型版本) + # 注意:同一模型不同参数规模可能有不同维度 + # 例如 qwen3-embedding:0.6b=1024, qwen3-embedding:8b=4096 + # 用户可通过 dimension 参数覆盖 MODELS = { "nomic-embed-text": 768, "mxbai-embed-large": 1024, "all-minilm": 384, "snowflake-arctic-embed": 1024, "bge-m3": 1024, - "qwen3-embedding": 1024, + "qwen3-embedding": 1024, # 默认值,8b版本为4096 } - + def __init__( self, base_url: Optional[str] = None, model: str = "nomic-embed-text", + dimension: Optional[int] = None, ): self.base_url = base_url or "http://localhost:11434" self.model = model - self._dimension = self.MODELS.get(model, 768) + # 用户指定的维度优先,否则使用默认映射 + self._dimension = dimension if dimension else self.MODELS.get(model, 768) @property def dimension(self) -> int: @@ -559,7 +565,7 @@ class EmbeddingService: """ 嵌入服务 统一管理嵌入模型和缓存 - + 支持的提供商: - openai: OpenAI 官方 - azure: Azure OpenAI @@ -568,13 +574,14 @@ class EmbeddingService: - huggingface: HuggingFace Inference API - jina: Jina AI """ - + def __init__( self, provider: Optional[str] = None, model: Optional[str] = None, api_key: Optional[str] = None, base_url: Optional[str] = None, + dimension: Optional[int] = None, cache_enabled: bool = True, ): """ @@ -585,6 +592,7 @@ class EmbeddingService: model: 模型名称 api_key: API Key base_url: API Base URL + dimension: 向量维度(可选,用于覆盖默认值) cache_enabled: 是否启用缓存 """ self.cache_enabled = cache_enabled @@ -595,6 +603,7 @@ class EmbeddingService: self.model = model or getattr(settings, 'EMBEDDING_MODEL', 'text-embedding-3-small') self.api_key = api_key self.base_url = base_url + self.custom_dimension = dimension # 创建提供商实例 self._provider = self._create_provider( @@ -602,9 +611,10 @@ class EmbeddingService: model=self.model, api_key=api_key, base_url=base_url, + dimension=dimension, ) - logger.info(f"Embedding service initialized with {self.provider}/{self.model}") + logger.info(f"Embedding service initialized with {self.provider}/{self.model}, dimension={self._provider.dimension}") def _create_provider( self, @@ -612,28 +622,29 @@ class EmbeddingService: model: str, api_key: Optional[str], base_url: Optional[str], + dimension: Optional[int] = None, ) -> EmbeddingProvider: """创建嵌入提供商实例""" provider = provider.lower() - + if provider == "ollama": - return OllamaEmbedding(base_url=base_url, model=model) - + return OllamaEmbedding(base_url=base_url, model=model, dimension=dimension) + elif provider == "azure": return AzureOpenAIEmbedding(api_key=api_key, base_url=base_url, model=model) - + elif provider == "cohere": return CohereEmbedding(api_key=api_key, base_url=base_url, model=model) - + elif provider == "huggingface": return HuggingFaceEmbedding(api_key=api_key, base_url=base_url, model=model) - + elif provider == "jina": return JinaEmbedding(api_key=api_key, base_url=base_url, model=model) - + elif provider == "qwen": return QwenEmbedding(api_key=api_key, base_url=base_url, model=model) - + else: # 默认使用 OpenAI return OpenAIEmbedding(api_key=api_key, base_url=base_url, model=model) diff --git a/frontend/src/components/agent/EmbeddingConfig.tsx b/frontend/src/components/agent/EmbeddingConfig.tsx index 564cc74..0852bfb 100644 --- a/frontend/src/components/agent/EmbeddingConfig.tsx +++ b/frontend/src/components/agent/EmbeddingConfig.tsx @@ -73,6 +73,7 @@ export default function EmbeddingConfigPanel() { const [selectedModel, setSelectedModel] = useState(""); const [apiKey, setApiKey] = useState(""); const [baseUrl, setBaseUrl] = useState(""); + const [customDimension, setCustomDimension] = useState(null); const [batchSize, setBatchSize] = useState(100); // 加载数据 @@ -107,6 +108,7 @@ export default function EmbeddingConfigPanel() { setSelectedModel(configRes.data.model); setApiKey(configRes.data.api_key || ""); setBaseUrl(configRes.data.base_url || ""); + setCustomDimension(configRes.data.dimensions || null); setBatchSize(configRes.data.batch_size); } } catch (error) { @@ -135,6 +137,7 @@ export default function EmbeddingConfigPanel() { model: selectedModel, api_key: apiKey || undefined, base_url: baseUrl || undefined, + dimensions: customDimension || undefined, batch_size: batchSize, }); @@ -162,6 +165,7 @@ export default function EmbeddingConfigPanel() { model: selectedModel, api_key: apiKey || undefined, base_url: baseUrl || undefined, + dimension: customDimension || undefined, }); setTestResult(response.data); @@ -340,6 +344,27 @@ export default function EmbeddingConfigPanel() {

+ {/* 自定义向量维度 */} +
+ + setCustomDimension(e.target.value ? parseInt(e.target.value) : null)} + placeholder="留空使用默认值" + min={64} + max={8192} + className="h-10 cyber-input w-40" + /> +

+ 适用于 Ollama 等场景:同一模型不同参数规模可能有不同维度 +
+ 例如 qwen3-embedding:0.6b=1024, qwen3-embedding:8b=4096 +

+
+ {/* 批处理大小 */}