fix(embedding): 传递本地索引限制并冻结推理配置
This commit is contained in:
@@ -31,6 +31,7 @@ class ParsedNote:
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
blocks: list[NoteBlock] = field(default_factory=list)
|
||||
embedding_local_only: bool = False
|
||||
|
||||
|
||||
def note_id_for_path(rel_path: str) -> str:
|
||||
@@ -69,6 +70,7 @@ def parse_note(
|
||||
created_at=created_at,
|
||||
updated_at=updated_at,
|
||||
blocks=blocks,
|
||||
embedding_local_only=str(frontmatter.get("embedding_local_only", "")).lower() == "true",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -165,21 +165,32 @@ runtime = Runtime()
|
||||
class LocalEmbedding:
|
||||
dim = 384
|
||||
|
||||
def __init__(self, config=None):
|
||||
self._config = config
|
||||
|
||||
def snapshot(self):
|
||||
return LocalEmbedding((self._config or configuration()).model_copy(deep=True))
|
||||
|
||||
@property
|
||||
def model_id(self):
|
||||
spec = CATALOG[configuration().embedding_model]
|
||||
spec = CATALOG[(self._config or configuration()).embedding_model]
|
||||
return f"{spec.repository}@{spec.revision}"
|
||||
|
||||
@property
|
||||
def version(self):
|
||||
return CATALOG[configuration().embedding_model].revision
|
||||
return CATALOG[(self._config or configuration()).embedding_model].revision
|
||||
|
||||
@property
|
||||
def available(self):
|
||||
return read_state(configuration().embedding_model)["status"] == "installed" and interpreter().is_file()
|
||||
|
||||
async def embed_documents(self, texts):
|
||||
return await runtime.infer(configuration().embedding_model, "embedding", {"texts": texts}, priority=0)
|
||||
config = (self._config or configuration()).model_copy(deep=True)
|
||||
token = runtime_context.set(config)
|
||||
try:
|
||||
return await runtime.infer(config.embedding_model, "embedding", {"texts": texts}, priority=0)
|
||||
finally:
|
||||
runtime_context.reset(token)
|
||||
|
||||
async def embed_query(self, query):
|
||||
return (await self.embed_documents([query]))[0]
|
||||
|
||||
@@ -206,9 +206,9 @@ class ModelRoutingService:
|
||||
raise invalid_response()
|
||||
return data, url
|
||||
|
||||
async def embed(self, texts: list[str]) -> EmbeddingResult:
|
||||
async def embed(self, texts: list[str], *, local_only=False) -> EmbeddingResult:
|
||||
config = self.configuration()
|
||||
binding = config.embedding
|
||||
binding = None if local_only else config.embedding
|
||||
record_embedding(route_version=config.version,
|
||||
requested_route=binding.model_dump() if binding else None)
|
||||
reason = None
|
||||
@@ -255,12 +255,14 @@ class ModelRoutingService:
|
||||
model_id="api-" + hashlib.sha256(identity.encode()).hexdigest())
|
||||
except ProviderError as exc:
|
||||
reason = exc.code
|
||||
from app.local_models.runtime import LocalEmbedding
|
||||
local_embedding = self.local_embedding.snapshot() if isinstance(self.local_embedding, LocalEmbedding) else self.local_embedding
|
||||
try:
|
||||
vectors = await self.local_embedding.embed_documents(texts)
|
||||
vectors = await local_embedding.embed_documents(texts)
|
||||
except ProviderError as exc:
|
||||
raise ApiError(503, exc.code, exc.message, {"fallback_reason": reason}) from exc
|
||||
return EmbeddingResult(vectors=vectors, source="local", model_id=self.local_embedding.model_id,
|
||||
dimensions=self.local_embedding.dim, fallback_reason=reason)
|
||||
return EmbeddingResult(vectors=vectors, source="local", model_id=local_embedding.model_id,
|
||||
dimensions=local_embedding.dim, fallback_reason=reason)
|
||||
|
||||
@staticmethod
|
||||
def _media_file(path: Path):
|
||||
|
||||
@@ -35,7 +35,7 @@ class EmbeddingResult(Protocol):
|
||||
|
||||
|
||||
class EmbeddingRuntime(Protocol):
|
||||
async def embed(self, texts: list[str]) -> EmbeddingResult: ...
|
||||
async def embed(self, texts: list[str], *, local_only=False) -> EmbeddingResult: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -69,7 +69,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], *, accept_local=False, strict=False) -> RemoteEmbeddings | None:
|
||||
async def embed_remote(texts: list[str], *, accept_local=False, strict=False, local_only=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
|
||||
@@ -83,7 +83,7 @@ async def embed_remote(texts: list[str], *, accept_local=False, strict=False) ->
|
||||
if strict:
|
||||
raise ApiError(503, "EMBEDDING_UNAVAILABLE", "Embedding 服务未就绪,请检查模型路由和本地运行环境。")
|
||||
return None
|
||||
result = await runtime.embed(texts)
|
||||
result = await runtime.embed(texts, local_only=True) if local_only else await runtime.embed(texts)
|
||||
if result.source != "api" and not accept_local:
|
||||
record_embedding(fallback_reason=result.fallback_reason)
|
||||
return None
|
||||
|
||||
@@ -41,6 +41,9 @@ async def create_transcript_note(job_id, options):
|
||||
lines.append("")
|
||||
else:
|
||||
lines.append(job.text or "")
|
||||
if job.local_only:
|
||||
# Persist the indexing policy in the Vault, including later rebuilds.
|
||||
lines = ["---", "embedding_local_only: true", "---", "", *lines]
|
||||
try:
|
||||
note = await note_service.create_note(title=title, markdown="\n".join(lines), folder=options.folder, tags=["转写"])
|
||||
except ApiError as exc:
|
||||
|
||||
@@ -82,10 +82,10 @@ async def prepare_note_index(parsed: ParsedNote, *, strict=False) -> PreparedInd
|
||||
texts = [block.content for block in parsed.blocks]
|
||||
if isinstance(embedding, LocalEmbedding):
|
||||
# One routed invocation: API first, validated local fallback. No hash vectors.
|
||||
remote = await routed_vectors.embed_remote(texts, accept_local=True, strict=strict)
|
||||
remote = await routed_vectors.embed_remote(texts, accept_local=True, strict=strict, local_only=parsed.embedding_local_only)
|
||||
return [], remote
|
||||
vectors = await embedding.embed_documents(texts)
|
||||
remote = await routed_vectors.embed_remote(texts)
|
||||
remote = await routed_vectors.embed_remote(texts, local_only=parsed.embedding_local_only)
|
||||
return vectors, remote
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user