embedflow 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- embedflow/__init__.py +25 -0
- embedflow/__main__.py +3 -0
- embedflow/analysis.py +192 -0
- embedflow/cache/__init__.py +4 -0
- embedflow/cache/base.py +28 -0
- embedflow/cache/persistent_cache.py +198 -0
- embedflow/cli.py +1200 -0
- embedflow/compatibility/__init__.py +28 -0
- embedflow/compatibility/candidate_gap.py +105 -0
- embedflow/compatibility/containment.py +17 -0
- embedflow/compatibility/evaluate.py +319 -0
- embedflow/compatibility/metrics.py +75 -0
- embedflow/compatibility/migration_depth.py +67 -0
- embedflow/compatibility/probe.py +34 -0
- embedflow/compatibility/report.py +102 -0
- embedflow/compatibility/t2.py +64 -0
- embedflow/config.py +455 -0
- embedflow/data/__init__.py +1 -0
- embedflow/data/registry/__init__.py +1 -0
- embedflow/data/registry/benchmark_profiles.jsonl +3 -0
- embedflow/data/registry/checksums.sha256 +4 -0
- embedflow/data/registry/migrations.jsonl +15 -0
- embedflow/data/registry/registry_manifest.json +16 -0
- embedflow/data/registry/research_summaries.json +55 -0
- embedflow/data/registry/schema_version.json +5 -0
- embedflow/frozen/T2_V1_FROZEN_SPEC.md +71 -0
- embedflow/frozen/T2_V1_FROZEN_SPEC.sha256 +1 -0
- embedflow/indexes/__init__.py +5 -0
- embedflow/indexes/base.py +60 -0
- embedflow/indexes/faiss_backend.py +240 -0
- embedflow/indexes/qdrant_backend.py +225 -0
- embedflow/metrics/__init__.py +3 -0
- embedflow/metrics/latency.py +50 -0
- embedflow/migration/__init__.py +3 -0
- embedflow/migration/compatibility.py +156 -0
- embedflow/migration/facade.py +312 -0
- embedflow/migration/materializer.py +190 -0
- embedflow/migration/planner.py +78 -0
- embedflow/migration/state.py +81 -0
- embedflow/models/__init__.py +4 -0
- embedflow/models/base.py +31 -0
- embedflow/models/huggingface.py +226 -0
- embedflow/registry/__init__.py +47 -0
- embedflow/registry/loader.py +785 -0
- embedflow/registry/matcher.py +197 -0
- embedflow/registry/schema.py +266 -0
- embedflow/runtime.py +115 -0
- embedflow/serving/__init__.py +3 -0
- embedflow/serving/api.py +161 -0
- embedflow/serving/engine.py +222 -0
- embedflow/serving/factory.py +3 -0
- embedflow/serving/schemas.py +39 -0
- embedflow-0.1.0.dist-info/METADATA +210 -0
- embedflow-0.1.0.dist-info/RECORD +64 -0
- embedflow-0.1.0.dist-info/WHEEL +5 -0
- embedflow-0.1.0.dist-info/entry_points.txt +2 -0
- embedflow-0.1.0.dist-info/licenses/LICENSE +178 -0
- embedflow-0.1.0.dist-info/top_level.txt +2 -0
- src/__init__.py +1 -0
- src/embed.py +123 -0
- src/probe_features.py +24 -0
- src/storage.py +51 -0
- src/t2_v1.py +21 -0
- src/utils.py +53 -0
|
@@ -0,0 +1,312 @@
|
|
|
1
|
+
"""Small public facade for starting a progressive model migration.
|
|
2
|
+
|
|
3
|
+
The research implementation deliberately exposes the lower-level config and
|
|
4
|
+
engine objects. This module adds the short path a user needs in an
|
|
5
|
+
application: point EmbedFlow at an existing FAISS/Qdrant index, name the old
|
|
6
|
+
and new embedding models, and receive a serving session. All retrieval,
|
|
7
|
+
cache, worker, and probe behavior is delegated to the existing implementation;
|
|
8
|
+
this is not a second query path.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
import json
|
|
14
|
+
from collections.abc import Iterable
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
from typing import Any
|
|
17
|
+
|
|
18
|
+
from ..cache import SQLiteVectorCache
|
|
19
|
+
from ..config import (
|
|
20
|
+
CacheConfig,
|
|
21
|
+
DocumentsConfig,
|
|
22
|
+
EmbedFlowConfig,
|
|
23
|
+
IndexConfig,
|
|
24
|
+
MigrationConfig,
|
|
25
|
+
ModelConfig,
|
|
26
|
+
hydrate_research_contract,
|
|
27
|
+
save_config,
|
|
28
|
+
)
|
|
29
|
+
from ..indexes.base import VectorIndex
|
|
30
|
+
from ..migration.compatibility import run_probe, save_probe
|
|
31
|
+
from ..migration.state import DocumentStore
|
|
32
|
+
from ..models import EmbeddingModel, load_embedding_model
|
|
33
|
+
from ..runtime import load_index
|
|
34
|
+
from ..serving.engine import MigrationEngine
|
|
35
|
+
|
|
36
|
+
_MODEL_REGISTRY = {
|
|
37
|
+
"sentence-transformers/all-MiniLM-L6-v2": "minilm_l6",
|
|
38
|
+
"Qwen/Qwen3-Embedding-0.6B": "qwen3_0_6b",
|
|
39
|
+
"Qwen/Qwen3-Embedding-4B": "qwen3_4b",
|
|
40
|
+
"Qwen/Qwen3-Embedding-8B": "qwen3_8b",
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _is_model(value: Any) -> bool:
|
|
45
|
+
return all(hasattr(value, name) for name in ("encode_queries", "encode_documents", "dimension", "fingerprint"))
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _model_config(value: str | EmbeddingModel, model_root: str | Path | None = None) -> ModelConfig:
|
|
49
|
+
"""Build a contract for a model identifier or an already-loaded model."""
|
|
50
|
+
if _is_model(value):
|
|
51
|
+
# HuggingFaceEmbeddingModel retains its exact hydrated config. Reuse
|
|
52
|
+
# it when available so fingerprints and prompt contracts remain exact.
|
|
53
|
+
retained = getattr(value, "cfg", None)
|
|
54
|
+
if isinstance(retained, ModelConfig):
|
|
55
|
+
return retained
|
|
56
|
+
return ModelConfig(model=str(value.model_id), dimension=int(value.dimension))
|
|
57
|
+
if not isinstance(value, (str, Path)):
|
|
58
|
+
raise TypeError("old_model/new_model must be a model ID/path or EmbeddingModel instance")
|
|
59
|
+
identifier = str(value)
|
|
60
|
+
cfg = hydrate_research_contract(ModelConfig(identifier), project_root=Path(__file__).resolve().parents[2])
|
|
61
|
+
candidate = Path(identifier)
|
|
62
|
+
if candidate.exists():
|
|
63
|
+
cfg.local_path = str(candidate.resolve())
|
|
64
|
+
elif model_root is not None and identifier in _MODEL_REGISTRY:
|
|
65
|
+
staged = Path(model_root).expanduser().resolve() / _MODEL_REGISTRY[identifier]
|
|
66
|
+
if staged.exists():
|
|
67
|
+
cfg.local_path = str(staged)
|
|
68
|
+
return cfg
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def _queries(value: str | Path | Iterable[tuple[str, str]]) -> list[tuple[str, str]]:
|
|
72
|
+
if isinstance(value, (str, Path)):
|
|
73
|
+
path = Path(value)
|
|
74
|
+
if not path.exists():
|
|
75
|
+
raise FileNotFoundError(path)
|
|
76
|
+
rows: list[tuple[str, str]] = []
|
|
77
|
+
with path.open() as handle:
|
|
78
|
+
seen: set[str] = set()
|
|
79
|
+
for i, line in enumerate(handle, 1):
|
|
80
|
+
if not line.strip():
|
|
81
|
+
continue
|
|
82
|
+
try:
|
|
83
|
+
row = json.loads(line)
|
|
84
|
+
except json.JSONDecodeError as exc:
|
|
85
|
+
raise ValueError(f"invalid query JSON at {path}:{i}") from exc
|
|
86
|
+
if not isinstance(row, dict):
|
|
87
|
+
raise ValueError(f"query row {i} in {path} must be a JSON object")
|
|
88
|
+
query_id = str(row.get("id", row.get("query_id", i)))
|
|
89
|
+
text = row.get("text", row.get("query"))
|
|
90
|
+
if query_id in seen:
|
|
91
|
+
raise ValueError(f"duplicate query ID {query_id!r} in {path}")
|
|
92
|
+
if not isinstance(text, str) or not text.strip():
|
|
93
|
+
raise ValueError(f"query {query_id!r} has no text")
|
|
94
|
+
seen.add(query_id)
|
|
95
|
+
rows.append((query_id, text))
|
|
96
|
+
if not rows:
|
|
97
|
+
raise ValueError(f"query file is empty: {value}")
|
|
98
|
+
return rows
|
|
99
|
+
try:
|
|
100
|
+
rows = [(str(query_id), text) for query_id, text in value]
|
|
101
|
+
except (TypeError, ValueError) as exc:
|
|
102
|
+
raise ValueError("probe_queries must be an iterable of (query_id, text) pairs") from exc
|
|
103
|
+
if len({query_id for query_id, _ in rows}) != len(rows):
|
|
104
|
+
raise ValueError("probe_queries must not contain duplicate query IDs")
|
|
105
|
+
if not rows or any(not isinstance(text, str) or not text.strip() for _, text in rows):
|
|
106
|
+
raise ValueError("probe_queries must contain non-empty query text")
|
|
107
|
+
return rows
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
class MigrationSession:
|
|
111
|
+
"""Handle returned by :func:`migrate`.
|
|
112
|
+
|
|
113
|
+
``search`` and ``status`` are synchronous convenience methods for an
|
|
114
|
+
application. ``serve`` mounts the same FastAPI dashboard/API used by the
|
|
115
|
+
CLI. Use ``close`` or a context manager to release models and the cache.
|
|
116
|
+
"""
|
|
117
|
+
|
|
118
|
+
def __init__(self, engine: MigrationEngine, config: EmbedFlowConfig,
|
|
119
|
+
config_path: Path | None, owns_source: bool, owns_target: bool, owns_index: bool):
|
|
120
|
+
self.engine = engine
|
|
121
|
+
self.config = config
|
|
122
|
+
self.config_path = config_path
|
|
123
|
+
self._owns_source = owns_source
|
|
124
|
+
self._owns_target = owns_target
|
|
125
|
+
self._owns_index = owns_index
|
|
126
|
+
self._closed = False
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def plan(self):
|
|
130
|
+
return self.engine.plan
|
|
131
|
+
|
|
132
|
+
def search(self, query: str, top_k: int = 10, candidate_depth: int | None = None,
|
|
133
|
+
max_sync_misses: int | None = None) -> dict[str, Any]:
|
|
134
|
+
return self.engine.search(query, top_k, candidate_depth, max_sync_misses)
|
|
135
|
+
|
|
136
|
+
def status(self) -> dict[str, Any]:
|
|
137
|
+
return self.engine.status()
|
|
138
|
+
|
|
139
|
+
def prewarm(self, document_ids: list[str], asynchronous: bool = True) -> dict[str, Any]:
|
|
140
|
+
return self.engine.prewarm(document_ids, asynchronous=asynchronous)
|
|
141
|
+
|
|
142
|
+
def app(self):
|
|
143
|
+
from ..serving.api import create_app
|
|
144
|
+
return create_app(self.engine)
|
|
145
|
+
|
|
146
|
+
def serve(self, host: str = "127.0.0.1", port: int = 8000, log_level: str = "info") -> None:
|
|
147
|
+
"""Run the API/dashboard until interrupted."""
|
|
148
|
+
try:
|
|
149
|
+
import uvicorn
|
|
150
|
+
except ImportError as exc:
|
|
151
|
+
raise RuntimeError("serving requires uvicorn; install `embedflow[dashboard]`") from exc
|
|
152
|
+
uvicorn.run(self.app(), host=host, port=int(port), log_level=log_level)
|
|
153
|
+
|
|
154
|
+
def close(self) -> None:
|
|
155
|
+
if self._closed:
|
|
156
|
+
return
|
|
157
|
+
self.engine.close(close_models=False, close_indexes=self._owns_index)
|
|
158
|
+
if self._owns_source:
|
|
159
|
+
self.engine.source_model.close()
|
|
160
|
+
if self._owns_target:
|
|
161
|
+
self.engine.target_model.close()
|
|
162
|
+
self._closed = True
|
|
163
|
+
|
|
164
|
+
def __enter__(self) -> MigrationSession:
|
|
165
|
+
return self
|
|
166
|
+
|
|
167
|
+
def __exit__(self, *_: Any) -> None:
|
|
168
|
+
self.close()
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
def migrate(
|
|
172
|
+
*,
|
|
173
|
+
index: str | Path | VectorIndex,
|
|
174
|
+
old_model: str | Path | EmbeddingModel,
|
|
175
|
+
new_model: str | Path | EmbeddingModel,
|
|
176
|
+
documents: str | Path | DocumentStore,
|
|
177
|
+
backend: str | None = None,
|
|
178
|
+
index_url: str | None = None,
|
|
179
|
+
collection: str = "embedflow",
|
|
180
|
+
vector_name: str | None = None,
|
|
181
|
+
api_key_env: str | None = "QDRANT_API_KEY",
|
|
182
|
+
metric: str = "cosine",
|
|
183
|
+
model_root: str | Path | None = None,
|
|
184
|
+
device: str | None = None,
|
|
185
|
+
cache_path: str | Path = "./embedflow_cache",
|
|
186
|
+
state_path: str | Path = "./embedflow_state.json",
|
|
187
|
+
config_path: str | Path | None = None,
|
|
188
|
+
candidate_depth: int = 50,
|
|
189
|
+
kmax_probe: int = 500,
|
|
190
|
+
max_sync_misses: int = 4,
|
|
191
|
+
background_batch_size: int = 32,
|
|
192
|
+
probe_queries: str | Path | Iterable[tuple[str, str]] | None = None,
|
|
193
|
+
probe_limit: int | None = None,
|
|
194
|
+
start_worker: bool = True,
|
|
195
|
+
) -> MigrationSession:
|
|
196
|
+
"""Start progressive migration over an existing index.
|
|
197
|
+
|
|
198
|
+
``index`` may be an existing EmbedFlow ``VectorIndex`` instance or a path
|
|
199
|
+
to a FAISS index (or a Qdrant path/URL when ``backend="qdrant"``). Model
|
|
200
|
+
strings are revision-hydrated when they are registered by the research
|
|
201
|
+
contract; model objects can be supplied by an application directly.
|
|
202
|
+
``probe_queries`` is optional so serving can start immediately. When
|
|
203
|
+
supplied, the existing frozen T2-v1 implementation is run before the
|
|
204
|
+
session is returned.
|
|
205
|
+
"""
|
|
206
|
+
if isinstance(documents, DocumentStore):
|
|
207
|
+
document_store = documents
|
|
208
|
+
else:
|
|
209
|
+
document_store = DocumentStore(str(documents))
|
|
210
|
+
source_cfg = _model_config(old_model, model_root)
|
|
211
|
+
target_cfg = _model_config(new_model, model_root)
|
|
212
|
+
selected_device = device or target_cfg.device or source_cfg.device or "cpu"
|
|
213
|
+
|
|
214
|
+
owns_source = not _is_model(old_model)
|
|
215
|
+
owns_target = not _is_model(new_model)
|
|
216
|
+
owns_index = not (isinstance(index, VectorIndex) or all(hasattr(index, name) for name in ("search", "size", "dimension", "metadata")))
|
|
217
|
+
source_model = None
|
|
218
|
+
target_model = None
|
|
219
|
+
source_index = None
|
|
220
|
+
cache = None
|
|
221
|
+
engine = None
|
|
222
|
+
try:
|
|
223
|
+
source_model = old_model if not owns_source else load_embedding_model(source_cfg, model_root=model_root, device=selected_device)
|
|
224
|
+
target_model = new_model if not owns_target else load_embedding_model(target_cfg, model_root=model_root, device=selected_device)
|
|
225
|
+
source_cfg.dimension = int(source_model.dimension)
|
|
226
|
+
target_cfg.dimension = int(target_model.dimension)
|
|
227
|
+
|
|
228
|
+
if not owns_index:
|
|
229
|
+
source_index = index
|
|
230
|
+
index_backend = backend or ("qdrant" if str(source_index.metadata().get("backend", "")).lower() == "qdrant" else "faiss")
|
|
231
|
+
index_path = str(getattr(source_index, "path", "./legacy.index") or "./legacy.index")
|
|
232
|
+
else:
|
|
233
|
+
index_path = str(index_url or index)
|
|
234
|
+
index_backend = (backend or ("qdrant" if "://" in index_path else "faiss")).lower()
|
|
235
|
+
if index_backend not in {"faiss", "qdrant"}:
|
|
236
|
+
raise ValueError("backend must be faiss or qdrant")
|
|
237
|
+
|
|
238
|
+
cfg = EmbedFlowConfig(
|
|
239
|
+
source=source_cfg,
|
|
240
|
+
target=target_cfg,
|
|
241
|
+
index=IndexConfig(backend=index_backend, path=index_path, collection=collection,
|
|
242
|
+
url=index_url, metric=str(metric), vector_name=vector_name,
|
|
243
|
+
api_key_env=api_key_env),
|
|
244
|
+
documents=DocumentsConfig(path=str(document_store.path)),
|
|
245
|
+
migration=MigrationConfig(candidate_depth=int(candidate_depth), kmax_probe=max(int(kmax_probe), int(candidate_depth)),
|
|
246
|
+
max_sync_misses=int(max_sync_misses), background_batch_size=int(background_batch_size)),
|
|
247
|
+
cache=CacheConfig(path=str(cache_path)),
|
|
248
|
+
state_path=str(state_path),
|
|
249
|
+
dashboard_title=f"EmbedFlow — {source_cfg.model} → {target_cfg.model}",
|
|
250
|
+
)
|
|
251
|
+
base = Path(config_path).expanduser().resolve().parent if config_path else Path.cwd()
|
|
252
|
+
cfg.resolve_paths(base)
|
|
253
|
+
cfg.validate()
|
|
254
|
+
|
|
255
|
+
if owns_index:
|
|
256
|
+
source_index = load_index(cfg, document_store)
|
|
257
|
+
if int(source_index.dimension) != int(source_model.dimension):
|
|
258
|
+
raise ValueError(f"source model dimension {source_model.dimension} != existing index dimension {source_index.dimension}")
|
|
259
|
+
stored_fingerprint = source_index.metadata().get("model_fingerprint")
|
|
260
|
+
if stored_fingerprint and stored_fingerprint != source_model.fingerprint:
|
|
261
|
+
raise ValueError("source model fingerprint does not match the existing index contract")
|
|
262
|
+
|
|
263
|
+
cache = SQLiteVectorCache(cfg.cache.path, target_model.fingerprint, target_model.dimension)
|
|
264
|
+
probe: dict[str, Any] = {}
|
|
265
|
+
if probe_queries is not None:
|
|
266
|
+
probe = run_probe(source_model, target_model, source_index, document_store, _queries(probe_queries),
|
|
267
|
+
kmax=cfg.migration.kmax_probe, seed=42, limit=probe_limit)
|
|
268
|
+
if config_path:
|
|
269
|
+
save_probe(probe, Path(config_path).expanduser().resolve().with_name("probe_result.json"))
|
|
270
|
+
if config_path:
|
|
271
|
+
config_file = Path(config_path).expanduser().resolve()
|
|
272
|
+
save_config(cfg, config_file)
|
|
273
|
+
else:
|
|
274
|
+
config_file = None
|
|
275
|
+
engine = MigrationEngine(cfg, source_model, target_model, source_index, cache, document_store,
|
|
276
|
+
probe=probe, start_worker=start_worker)
|
|
277
|
+
return MigrationSession(engine, cfg, config_file, owns_source, owns_target, owns_index)
|
|
278
|
+
except Exception:
|
|
279
|
+
# A failed setup must not leak a loaded model, SQLite connection, or a
|
|
280
|
+
# local Qdrant file lock. Only close resources created by this call;
|
|
281
|
+
# caller-owned model/index objects remain their caller's responsibility.
|
|
282
|
+
if engine is not None:
|
|
283
|
+
try:
|
|
284
|
+
engine.close(close_models=False, close_indexes=owns_index)
|
|
285
|
+
except Exception:
|
|
286
|
+
pass
|
|
287
|
+
elif cache is not None:
|
|
288
|
+
try:
|
|
289
|
+
cache.close()
|
|
290
|
+
except Exception:
|
|
291
|
+
pass
|
|
292
|
+
if owns_index and source_index is not None:
|
|
293
|
+
try:
|
|
294
|
+
close_index = getattr(source_index, "close", None)
|
|
295
|
+
if callable(close_index):
|
|
296
|
+
close_index()
|
|
297
|
+
except Exception:
|
|
298
|
+
pass
|
|
299
|
+
if owns_source and source_model is not None:
|
|
300
|
+
try:
|
|
301
|
+
source_model.close()
|
|
302
|
+
except Exception:
|
|
303
|
+
pass
|
|
304
|
+
if owns_target and target_model is not None:
|
|
305
|
+
try:
|
|
306
|
+
target_model.close()
|
|
307
|
+
except Exception:
|
|
308
|
+
pass
|
|
309
|
+
raise
|
|
310
|
+
|
|
311
|
+
|
|
312
|
+
__all__ = ["MigrationSession", "migrate"]
|
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import sqlite3
|
|
4
|
+
import threading
|
|
5
|
+
import time
|
|
6
|
+
from collections.abc import Iterable
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
from ..cache import SQLiteVectorCache
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class PersistentWorkQueue:
|
|
16
|
+
def __init__(self, path: str | Path, max_retries: int = 3):
|
|
17
|
+
p = Path(path); self.path = p if p.suffix else p / "materialization_queue.sqlite3"
|
|
18
|
+
try:
|
|
19
|
+
retries = int(max_retries)
|
|
20
|
+
exact = float(max_retries) == retries
|
|
21
|
+
except (TypeError, ValueError, OverflowError):
|
|
22
|
+
retries, exact = 0, False
|
|
23
|
+
if isinstance(max_retries, bool) or not exact or retries < 1:
|
|
24
|
+
raise ValueError("max_retries must be a positive integer")
|
|
25
|
+
self.path.parent.mkdir(parents=True, exist_ok=True); self.max_retries = retries
|
|
26
|
+
self.db = None; self.lock = threading.RLock()
|
|
27
|
+
try:
|
|
28
|
+
self.db = sqlite3.connect(str(self.path), check_same_thread=False, timeout=30)
|
|
29
|
+
self.db.execute("PRAGMA busy_timeout=30000")
|
|
30
|
+
self.db.execute("PRAGMA journal_mode=WAL")
|
|
31
|
+
self.db.execute("""CREATE TABLE IF NOT EXISTS work (
|
|
32
|
+
document_id TEXT PRIMARY KEY, status TEXT NOT NULL, attempts INTEGER NOT NULL DEFAULT 0,
|
|
33
|
+
error TEXT, enqueued_at REAL NOT NULL, updated_at REAL NOT NULL)""")
|
|
34
|
+
self.db.execute("UPDATE work SET status='pending',updated_at=? WHERE status='processing'", (time.time(),)); self.db.commit()
|
|
35
|
+
except sqlite3.DatabaseError as exc:
|
|
36
|
+
if self.db is not None:
|
|
37
|
+
self.db.close()
|
|
38
|
+
self.db = None
|
|
39
|
+
raise ValueError(f"invalid materialization queue database: {self.path}") from exc
|
|
40
|
+
|
|
41
|
+
def _ensure_open(self) -> None:
|
|
42
|
+
if self.db is None:
|
|
43
|
+
raise RuntimeError("materialization queue is closed")
|
|
44
|
+
|
|
45
|
+
def enqueue(self, ids: Iterable[str]) -> int:
|
|
46
|
+
now = time.time(); n = 0
|
|
47
|
+
with self.lock:
|
|
48
|
+
self._ensure_open()
|
|
49
|
+
for document_id in dict.fromkeys(str(x) for x in ids):
|
|
50
|
+
cur = self.db.execute("INSERT OR IGNORE INTO work(document_id,status,enqueued_at,updated_at) VALUES(?,?,?,?)",
|
|
51
|
+
(document_id, "pending", now, now)); n += cur.rowcount
|
|
52
|
+
self.db.commit()
|
|
53
|
+
return n
|
|
54
|
+
|
|
55
|
+
def claim(self, limit: int) -> list[str]:
|
|
56
|
+
if isinstance(limit, bool) or int(limit) != limit or int(limit) < 1:
|
|
57
|
+
raise ValueError("queue claim limit must be a positive integer")
|
|
58
|
+
with self.lock:
|
|
59
|
+
self._ensure_open()
|
|
60
|
+
# A write transaction makes SELECT+UPDATE atomic across multiple
|
|
61
|
+
# worker processes sharing the same SQLite queue. Without it,
|
|
62
|
+
# two workers can claim the same pending document before either
|
|
63
|
+
# marks it processing.
|
|
64
|
+
self.db.execute("BEGIN IMMEDIATE")
|
|
65
|
+
try:
|
|
66
|
+
rows = self.db.execute("SELECT document_id FROM work WHERE status='pending' AND attempts<? ORDER BY enqueued_at,document_id LIMIT ?",
|
|
67
|
+
(self.max_retries, int(limit))).fetchall()
|
|
68
|
+
ids = [str(x[0]) for x in rows]; now = time.time()
|
|
69
|
+
self.db.executemany("UPDATE work SET status='processing',attempts=attempts+1,updated_at=? WHERE document_id=?",
|
|
70
|
+
[(now, x) for x in ids]); self.db.commit(); return ids
|
|
71
|
+
except Exception:
|
|
72
|
+
self.db.rollback()
|
|
73
|
+
raise
|
|
74
|
+
|
|
75
|
+
def complete(self, ids: Iterable[str]) -> None:
|
|
76
|
+
with self.lock:
|
|
77
|
+
self._ensure_open()
|
|
78
|
+
self.db.executemany("UPDATE work SET status='done',updated_at=? WHERE document_id=?", [(time.time(), str(x)) for x in ids]); self.db.commit()
|
|
79
|
+
|
|
80
|
+
def fail(self, ids: Iterable[str], error: str) -> None:
|
|
81
|
+
if not str(error).strip():
|
|
82
|
+
error = "unknown materialization failure"
|
|
83
|
+
with self.lock:
|
|
84
|
+
self._ensure_open()
|
|
85
|
+
self.db.executemany("UPDATE work SET status=CASE WHEN attempts<? THEN 'pending' ELSE 'error' END,error=?,updated_at=? WHERE document_id=?",
|
|
86
|
+
[(self.max_retries, str(error), time.time(), str(x)) for x in ids]); self.db.commit()
|
|
87
|
+
|
|
88
|
+
def stats(self) -> dict[str, int]:
|
|
89
|
+
with self.lock:
|
|
90
|
+
self._ensure_open()
|
|
91
|
+
rows = self.db.execute("SELECT status,COUNT(*) FROM work GROUP BY status").fetchall()
|
|
92
|
+
out = {"pending": 0, "processing": 0, "done": 0, "error": 0}; out.update({str(k): int(v) for k, v in rows}); return out
|
|
93
|
+
|
|
94
|
+
def close(self) -> None:
|
|
95
|
+
with self.lock:
|
|
96
|
+
if self.db is not None:
|
|
97
|
+
self.db.close()
|
|
98
|
+
self.db = None
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
class MaterializationWorker:
|
|
102
|
+
def __init__(self, target_model: Any, documents: Any, cache: SQLiteVectorCache,
|
|
103
|
+
queue_path: str | Path, batch_size: int = 32, max_retries: int = 3,
|
|
104
|
+
state: Any | None = None):
|
|
105
|
+
self.target_model, self.documents, self.cache = target_model, documents, cache
|
|
106
|
+
try:
|
|
107
|
+
parsed_batch_size = int(batch_size)
|
|
108
|
+
exact = float(batch_size) == parsed_batch_size
|
|
109
|
+
except (TypeError, ValueError, OverflowError):
|
|
110
|
+
parsed_batch_size, exact = 0, False
|
|
111
|
+
if isinstance(batch_size, bool) or not exact or parsed_batch_size < 1:
|
|
112
|
+
raise ValueError("batch_size must be a positive integer")
|
|
113
|
+
self.queue = PersistentWorkQueue(queue_path, max_retries); self.batch_size = parsed_batch_size; self.state = state
|
|
114
|
+
self.stop_event = threading.Event(); self.thread: threading.Thread | None = None
|
|
115
|
+
self._closed = False
|
|
116
|
+
self._stats = {"materialized": 0, "errors": 0, "started_at": None, "last_throughput_docs_sec": 0.0}
|
|
117
|
+
|
|
118
|
+
def enqueue(self, ids: Iterable[str]) -> int:
|
|
119
|
+
if self._closed:
|
|
120
|
+
raise RuntimeError("materialization worker is closed")
|
|
121
|
+
return self.queue.enqueue(ids)
|
|
122
|
+
|
|
123
|
+
def start(self) -> None:
|
|
124
|
+
if self._closed:
|
|
125
|
+
raise RuntimeError("materialization worker is closed")
|
|
126
|
+
if self.thread and self.thread.is_alive(): return
|
|
127
|
+
self.stop_event.clear(); self._stats["started_at"] = time.time(); self.thread = threading.Thread(target=self._run, name="embedflow-materializer", daemon=True); self.thread.start()
|
|
128
|
+
|
|
129
|
+
def _run(self) -> None:
|
|
130
|
+
while not self.stop_event.is_set():
|
|
131
|
+
try:
|
|
132
|
+
ids = self.queue.claim(self.batch_size)
|
|
133
|
+
except RuntimeError:
|
|
134
|
+
if self.stop_event.is_set() or self._closed:
|
|
135
|
+
return
|
|
136
|
+
raise
|
|
137
|
+
except Exception as exc:
|
|
138
|
+
self._stats["errors"] += 1
|
|
139
|
+
if self.state:
|
|
140
|
+
self.state.add_error(f"materialization queue failed: {exc}")
|
|
141
|
+
time.sleep(0.1)
|
|
142
|
+
continue
|
|
143
|
+
if not ids:
|
|
144
|
+
time.sleep(0.05)
|
|
145
|
+
continue
|
|
146
|
+
try:
|
|
147
|
+
# Resolve documents individually so one missing/deleted row is
|
|
148
|
+
# marked as an error without poisoning the rest of a batch.
|
|
149
|
+
available: list[str] = []
|
|
150
|
+
texts: list[str] = []
|
|
151
|
+
missing: list[str] = []
|
|
152
|
+
for document_id in ids:
|
|
153
|
+
try:
|
|
154
|
+
value = self.documents.get([document_id]).get(document_id)
|
|
155
|
+
except Exception:
|
|
156
|
+
value = None
|
|
157
|
+
if value is None:
|
|
158
|
+
missing.append(document_id)
|
|
159
|
+
else:
|
|
160
|
+
available.append(document_id); texts.append(str(value))
|
|
161
|
+
if missing:
|
|
162
|
+
self.queue.fail(missing, "document text is unavailable")
|
|
163
|
+
self._stats["errors"] += len(missing)
|
|
164
|
+
if not available:
|
|
165
|
+
continue
|
|
166
|
+
vectors = self.target_model.encode_documents(texts, batch_size=self.batch_size)
|
|
167
|
+
self.cache.put(available, np.asarray(vectors, dtype="float32")); self.queue.complete(available)
|
|
168
|
+
self._stats["materialized"] += len(available)
|
|
169
|
+
elapsed = max(time.time() - float(self._stats["started_at"] or time.time()), 1e-6)
|
|
170
|
+
self._stats["last_throughput_docs_sec"] = self._stats["materialized"] / elapsed
|
|
171
|
+
if self.state: self.state.update(materialized_documents=self._stats["materialized"], materializer_throughput_docs_sec=self._stats["last_throughput_docs_sec"], queue=self.queue.stats())
|
|
172
|
+
except Exception as exc:
|
|
173
|
+
self.queue.fail(ids, repr(exc)); self._stats["errors"] += len(ids)
|
|
174
|
+
if self.state: self.state.add_error(f"materialization failed for {len(ids)} docs: {exc}")
|
|
175
|
+
|
|
176
|
+
def stop(self, timeout: float = 5.0) -> None:
|
|
177
|
+
self.stop_event.set()
|
|
178
|
+
if self.thread:
|
|
179
|
+
self.thread.join(timeout=timeout)
|
|
180
|
+
if self.thread.is_alive():
|
|
181
|
+
raise RuntimeError("materialization worker did not stop before timeout")
|
|
182
|
+
|
|
183
|
+
def stats(self) -> dict[str, Any]: return {**self._stats, "queue": self.queue.stats()}
|
|
184
|
+
|
|
185
|
+
def close(self, timeout: float = 5.0) -> None:
|
|
186
|
+
if self._closed:
|
|
187
|
+
return
|
|
188
|
+
self.stop(timeout=timeout)
|
|
189
|
+
self.queue.close()
|
|
190
|
+
self._closed = True
|
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import asdict, dataclass
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@dataclass
|
|
8
|
+
class MigrationPlan:
|
|
9
|
+
status: str
|
|
10
|
+
target_model: str
|
|
11
|
+
candidate_depth: int
|
|
12
|
+
diagnostic: str
|
|
13
|
+
ann_status: str
|
|
14
|
+
corpus_size: int
|
|
15
|
+
cached_target_vectors: int
|
|
16
|
+
cache_fraction: float
|
|
17
|
+
migration_strategy: str
|
|
18
|
+
warnings: list[str]
|
|
19
|
+
rationale: str = ""
|
|
20
|
+
source_model: str = ""
|
|
21
|
+
|
|
22
|
+
def to_dict(self) -> dict[str, Any]: return asdict(self)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def make_plan(*, target_model: str, corpus_size: int, diagnostic: str,
|
|
26
|
+
recommended_k: int | None, ann_status: str, cached_target_vectors: int,
|
|
27
|
+
default_k: int = 50, throughput_docs_per_second: float | None = None,
|
|
28
|
+
gpu_price_per_hour: float | None = None, source_model: str = "") -> MigrationPlan:
|
|
29
|
+
diagnostic = str(diagnostic or "UNKNOWN").upper()
|
|
30
|
+
if diagnostic not in {"SAFE", "EXPAND", "UNSAFE_OR_UNCERTAIN", "UNKNOWN"}:
|
|
31
|
+
raise ValueError(f"unsupported diagnostic: {diagnostic}")
|
|
32
|
+
ann_value = str(ann_status or "UNKNOWN").upper()
|
|
33
|
+
if ann_value not in {"PASS", "WARNING", "UNKNOWN"}:
|
|
34
|
+
raise ValueError(f"unsupported ANN status: {ann_status}")
|
|
35
|
+
try:
|
|
36
|
+
corpus_int = int(corpus_size)
|
|
37
|
+
corpus_exact = float(corpus_size) == corpus_int
|
|
38
|
+
cached_int = int(cached_target_vectors)
|
|
39
|
+
cached_exact = float(cached_target_vectors) == cached_int
|
|
40
|
+
except (TypeError, ValueError, OverflowError) as exc:
|
|
41
|
+
raise ValueError("corpus_size and cached_target_vectors must be integers") from exc
|
|
42
|
+
if isinstance(corpus_size, bool) or not corpus_exact or corpus_int < 0:
|
|
43
|
+
raise ValueError("corpus_size must be a non-negative integer")
|
|
44
|
+
if isinstance(cached_target_vectors, bool) or not cached_exact or cached_int < 0:
|
|
45
|
+
raise ValueError("cached_target_vectors must be a non-negative integer")
|
|
46
|
+
# ``candidate_depth: auto`` is a configuration convenience. The actual
|
|
47
|
+
# serving plan always resolves to an integer, using the probe recommendation
|
|
48
|
+
# when available and the conservative default otherwise.
|
|
49
|
+
if str(default_k).lower() == "auto":
|
|
50
|
+
fallback_k = 50
|
|
51
|
+
else:
|
|
52
|
+
try:
|
|
53
|
+
fallback_k = int(default_k)
|
|
54
|
+
if isinstance(default_k, bool) or float(default_k) != fallback_k or fallback_k < 1:
|
|
55
|
+
raise ValueError
|
|
56
|
+
except (TypeError, ValueError, OverflowError) as exc:
|
|
57
|
+
raise ValueError("default_k must be auto or a positive integer") from exc
|
|
58
|
+
raw_k = fallback_k if recommended_k is None else recommended_k
|
|
59
|
+
try:
|
|
60
|
+
k = int(raw_k)
|
|
61
|
+
exact_k = float(raw_k) == k
|
|
62
|
+
except (TypeError, ValueError, OverflowError) as exc:
|
|
63
|
+
raise ValueError("recommended_k must be a positive integer") from exc
|
|
64
|
+
if isinstance(raw_k, bool) or not exact_k or k < 1:
|
|
65
|
+
raise ValueError("recommended_k must be a positive integer")
|
|
66
|
+
warnings: list[str] = []
|
|
67
|
+
if diagnostic != "SAFE":
|
|
68
|
+
warnings.append("Compatibility is not a SAFE T2-v1 diagnostic; validate manually before relying on progressive migration.")
|
|
69
|
+
if ann_value in {"WARNING", "UNKNOWN"}:
|
|
70
|
+
warnings.append("ANN fidelity is not confirmed; T2-v1 does not diagnose ANN approximation error.")
|
|
71
|
+
if cached_int > corpus_int:
|
|
72
|
+
warnings.append("Cache count exceeds corpus size; check document identity and cache contract.")
|
|
73
|
+
fraction = (cached_int / corpus_int) if corpus_int else 0.0
|
|
74
|
+
strategy = "PROGRESSIVE" if diagnostic == "SAFE" else "PROGRESSIVE_WITH_REVIEW"
|
|
75
|
+
status = "READY" if diagnostic == "SAFE" and ann_value in {"PASS", "UNKNOWN"} else "REVIEW"
|
|
76
|
+
rationale = "Finite-pool diagnostic supports the selected starting K; this is an empirical recommendation, not a guarantee." if diagnostic == "SAFE" else "Use source-only/full-backfill fallback until compatibility and ANN health are reviewed."
|
|
77
|
+
return MigrationPlan(status, target_model, k, diagnostic, ann_value, corpus_int,
|
|
78
|
+
cached_int, float(fraction), strategy, warnings, rationale, source_model)
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import threading
|
|
5
|
+
from collections.abc import Iterable
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
from ..config import EmbedFlowConfig
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class DocumentStore:
|
|
13
|
+
"""Memory-indexed JSONL document store for progressive materialization."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, path: str | Path, id_field: str = "id", text_field: str = "text"):
|
|
16
|
+
self.path = Path(path)
|
|
17
|
+
if not self.path.exists(): raise FileNotFoundError(self.path)
|
|
18
|
+
if not isinstance(id_field, str) or not id_field.strip() or not isinstance(text_field, str) or not text_field.strip():
|
|
19
|
+
raise ValueError("document id_field and text_field must be non-empty strings")
|
|
20
|
+
self.documents: dict[str, str] = {}
|
|
21
|
+
self.id_field, self.text_field = id_field, text_field
|
|
22
|
+
with self.path.open() as handle:
|
|
23
|
+
for line_no, line in enumerate(handle, 1):
|
|
24
|
+
if not line.strip(): continue
|
|
25
|
+
try: row = json.loads(line)
|
|
26
|
+
except json.JSONDecodeError as exc: raise ValueError(f"invalid JSON at {self.path}:{line_no}") from exc
|
|
27
|
+
if not isinstance(row, dict): raise ValueError(f"document row {line_no} in {self.path} must be a JSON object")
|
|
28
|
+
if id_field not in row or text_field not in row: raise ValueError(f"missing {id_field}/{text_field} at {self.path}:{line_no}")
|
|
29
|
+
document_id = str(row[id_field])
|
|
30
|
+
if not document_id.strip(): raise ValueError(f"empty document ID at {self.path}:{line_no}")
|
|
31
|
+
if document_id in self.documents: raise ValueError(f"duplicate document ID {document_id!r}")
|
|
32
|
+
text = row[text_field]
|
|
33
|
+
if text is None or not str(text).strip(): raise ValueError(f"missing document text for {document_id!r}")
|
|
34
|
+
self.documents[document_id] = str(text)
|
|
35
|
+
if not self.documents: raise ValueError(f"document file is empty: {self.path}")
|
|
36
|
+
|
|
37
|
+
def get(self, document_ids: Iterable[str]) -> dict[str, str]:
|
|
38
|
+
ids = [str(x) for x in document_ids]
|
|
39
|
+
missing = [x for x in ids if x not in self.documents]
|
|
40
|
+
if missing: raise KeyError(f"document text missing for IDs: {missing[:5]}")
|
|
41
|
+
return {x: self.documents[x] for x in ids}
|
|
42
|
+
|
|
43
|
+
def size(self) -> int: return len(self.documents)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class MigrationState:
|
|
47
|
+
def __init__(self, path: str | Path, cfg: EmbedFlowConfig):
|
|
48
|
+
self.path = Path(path); self._lock = threading.RLock()
|
|
49
|
+
self.data: dict[str, Any] = {
|
|
50
|
+
"source_model": cfg.source.model, "target_model": cfg.target.model,
|
|
51
|
+
"target_model_fingerprint": cfg.target.fingerprint, "status": "INITIALIZING",
|
|
52
|
+
"diagnostic": "UNKNOWN", "ann_status": "UNKNOWN", "candidate_depth": cfg.migration.candidate_depth,
|
|
53
|
+
"migration_strategy": "PROGRESSIVE", "queries": 0, "errors": [],
|
|
54
|
+
}
|
|
55
|
+
if self.path.exists():
|
|
56
|
+
try:
|
|
57
|
+
loaded = json.loads(self.path.read_text())
|
|
58
|
+
except (OSError, json.JSONDecodeError) as exc: raise ValueError(f"invalid state file {self.path}") from exc
|
|
59
|
+
if not isinstance(loaded, dict):
|
|
60
|
+
raise ValueError(f"state file {self.path} must contain a JSON object")
|
|
61
|
+
stored_fingerprint = loaded.get("target_model_fingerprint")
|
|
62
|
+
if stored_fingerprint and str(stored_fingerprint) != str(cfg.target.fingerprint):
|
|
63
|
+
raise ValueError("state file target model fingerprint does not match the configured target contract")
|
|
64
|
+
self.data.update(loaded)
|
|
65
|
+
self.save()
|
|
66
|
+
|
|
67
|
+
def update(self, **values: Any) -> None:
|
|
68
|
+
with self._lock:
|
|
69
|
+
self.data.update(values); self.save()
|
|
70
|
+
|
|
71
|
+
def add_error(self, message: str) -> None:
|
|
72
|
+
with self._lock:
|
|
73
|
+
errors = self.data.setdefault("errors", []); errors.append(str(message)); self.data["errors"] = errors[-100:]; self.save()
|
|
74
|
+
|
|
75
|
+
def save(self) -> None:
|
|
76
|
+
self.path.parent.mkdir(parents=True, exist_ok=True)
|
|
77
|
+
tmp = self.path.with_suffix(self.path.suffix + ".tmp")
|
|
78
|
+
tmp.write_text(json.dumps(self.data, indent=2, sort_keys=True, ensure_ascii=False)); tmp.replace(self.path)
|
|
79
|
+
|
|
80
|
+
def snapshot(self) -> dict[str, Any]:
|
|
81
|
+
with self._lock: return json.loads(json.dumps(self.data))
|