mirror of
https://github.com/lintsinghua/DeepAudit.git
synced 2026-08-30 17:20:05 +08:00
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:
@@ -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)
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
Reference in New Issue
Block a user