fix: 嵌入模型向量维度可配置,解决 Ollama 不同参数规模模型维度不匹配问题

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 <noreply@anthropic.com>
This commit is contained in:
lintsinghua
2026-01-24 14:42:31 +08:00
parent da853fdd8c
commit 8198f2f390
3 changed files with 71 additions and 27 deletions
@@ -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)
+27 -16
View File
@@ -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)
@@ -73,6 +73,7 @@ export default function EmbeddingConfigPanel() {
const [selectedModel, setSelectedModel] = useState("");
const [apiKey, setApiKey] = useState("");
const [baseUrl, setBaseUrl] = useState("");
const [customDimension, setCustomDimension] = useState<number | null>(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() {
</p>
</div>
{/* 自定义向量维度 */}
<div className="space-y-2">
<Label className="text-xs font-bold text-muted-foreground uppercase">
<span className="text-muted-foreground">()</span>
</Label>
<Input
type="number"
value={customDimension || ""}
onChange={(e) => setCustomDimension(e.target.value ? parseInt(e.target.value) : null)}
placeholder="留空使用默认值"
min={64}
max={8192}
className="h-10 cyber-input w-40"
/>
<p className="text-xs text-muted-foreground">
Ollama
<br />
qwen3-embedding:0.6b=1024, qwen3-embedding:8b=4096
</p>
</div>
{/* 批处理大小 */}
<div className="space-y-2">
<Label className="text-xs font-bold text-muted-foreground uppercase"></Label>