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
+
+
+
{/* 批处理大小 */}