feat(multimodal): 实现本地模型管线与请求用量配置
This commit is contained in:
@@ -1,7 +1,6 @@
|
||||
"""Embedding 统一接口与轻量实现。
|
||||
|
||||
真实默认是本地 BGE-M3 类模型,但第一阶段先跑通链路,这里用确定性的特征哈希向量代替。
|
||||
后续接入真实模型时实现同样的 EmbeddingProvider 接口替换即可,上层检索逻辑不变。
|
||||
生产环境使用 local_models 的真实模型。特征哈希实现仅供测试显式注入。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -20,6 +20,7 @@ from app.contracts import (
|
||||
)
|
||||
from app.repository import BlockHit
|
||||
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
|
||||
from app.local_models.runtime import LocalEmbedding
|
||||
from app.retrieval.hybrid import normalize_scores, rrf_fuse
|
||||
from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider
|
||||
from app.retrieval import routed_vectors
|
||||
@@ -88,8 +89,13 @@ class RetrievalEngine:
|
||||
and self.embedding is self._routed_defaults[0]
|
||||
and self.vector_store is self._routed_defaults[1]
|
||||
):
|
||||
vec_hits = await routed_vectors.search_remote(request.query, top_k=recall)
|
||||
vec_hits = await routed_vectors.search_remote(request.query, top_k=recall, accept_local=isinstance(self.embedding, LocalEmbedding))
|
||||
if vec_hits is None:
|
||||
if isinstance(self.embedding, LocalEmbedding):
|
||||
if request.mode == SearchMode.hybrid:
|
||||
return self._search_fts(request)
|
||||
from app.errors import ApiError
|
||||
raise ApiError(409, "SEMANTIC_INDEX_UNAVAILABLE", "语义索引未就绪。请配置 Embedding 或下载本地模型后重建索引。")
|
||||
query_vec = await self.embedding.embed_query(request.query)
|
||||
vec_hits = await self.vector_store.search(query_vec, top_k=recall)
|
||||
record_embedding(source="local", model_id=self.embedding.model_id,
|
||||
@@ -287,5 +293,5 @@ def _utc(dt: datetime) -> datetime:
|
||||
|
||||
# 默认引擎实例:轻量实现跑通链路,后续可替换真实模型实现
|
||||
engine = RetrievalEngine(
|
||||
HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore(), route_embeddings=True,
|
||||
LocalEmbedding(), LexicalReranker(), SqliteVecStore(), route_embeddings=True,
|
||||
)
|
||||
|
||||
@@ -42,6 +42,7 @@ class RemoteEmbeddings:
|
||||
space_id: str
|
||||
dimensions: int
|
||||
vectors: list[list[float]]
|
||||
source: str = "api"
|
||||
|
||||
|
||||
def get_model_routing() -> EmbeddingRuntime | None:
|
||||
@@ -67,7 +68,7 @@ def _unit_vector(vector: list[float], dimensions: int) -> list[float]:
|
||||
return [value / norm for value in scaled]
|
||||
|
||||
|
||||
async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
async def embed_remote(texts: list[str], *, accept_local=False) -> RemoteEmbeddings | None:
|
||||
"""Return validated API vectors, or None to use the caller's local baseline.
|
||||
|
||||
Do not use the runtime's local result: the caller may have injected its own
|
||||
@@ -80,7 +81,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
if runtime is None:
|
||||
return None
|
||||
result = await runtime.embed(texts)
|
||||
if result.source != "api":
|
||||
if result.source != "api" and not accept_local:
|
||||
record_embedding(fallback_reason=result.fallback_reason)
|
||||
return None
|
||||
if not isinstance(result.model_id, str) or not result.model_id or result.model_id == "hash-v1":
|
||||
@@ -93,6 +94,7 @@ async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
space_id=result.model_id,
|
||||
dimensions=result.dimensions,
|
||||
vectors=[_unit_vector(vector, result.dimensions) for vector in result.vectors],
|
||||
source=result.source,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Avoid logging provider exceptions containing credentials or note text.
|
||||
@@ -152,13 +154,13 @@ def store_remote(
|
||||
logger.warning("Remote vector storage unavailable (%s); local index retained", type(exc).__name__)
|
||||
|
||||
|
||||
async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
|
||||
async def search_remote(query: str, *, top_k: int, accept_local=False) -> list[VectorHit] | None:
|
||||
"""None means fallback, including any missing/invalid current-block vector.
|
||||
|
||||
Read coverage and vectors together so concurrent note updates cannot produce
|
||||
an apparently complete subset. Never fill missing remote hits with local hits.
|
||||
"""
|
||||
batch = await embed_remote([query])
|
||||
batch = await embed_remote([query], accept_local=accept_local)
|
||||
if batch is None:
|
||||
return None
|
||||
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
|
||||
@@ -190,7 +192,7 @@ async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
|
||||
yield VectorHit(id=row["block_id"], score=max(0.0, min(1.0, score)))
|
||||
|
||||
result = heapq.nlargest(top_k, hits(), key=lambda hit: hit.score)
|
||||
record_embedding(source="api", model_id=batch.space_id,
|
||||
record_embedding(source=batch.source, model_id=batch.space_id,
|
||||
dimensions=batch.dimensions, fallback_reason=None)
|
||||
return result
|
||||
finally:
|
||||
|
||||
Reference in New Issue
Block a user