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,225 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import uuid
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
from .base import SearchHit, VectorIndex, validate_k, validate_query_vector
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class QdrantIndex(VectorIndex):
|
|
13
|
+
"""Qdrant adapter. The dependency is optional and imported lazily."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, client: Any, collection: str, dimension: int,
|
|
16
|
+
documents: dict[str, Any] | None = None, metric: str = "cosine",
|
|
17
|
+
vector_name: str | None = None):
|
|
18
|
+
try:
|
|
19
|
+
dimension_value = int(dimension)
|
|
20
|
+
dimension_exact = float(dimension) == dimension_value
|
|
21
|
+
except (TypeError, ValueError, OverflowError) as exc:
|
|
22
|
+
raise ValueError("Qdrant vector dimension must be a positive integer") from exc
|
|
23
|
+
if isinstance(dimension, bool) or not dimension_exact or dimension_value < 1:
|
|
24
|
+
raise ValueError("Qdrant vector dimension must be a positive integer")
|
|
25
|
+
metric = str(metric).lower()
|
|
26
|
+
if metric not in {"cosine", "dot", "inner_product"}:
|
|
27
|
+
raise ValueError("metric must be cosine, dot, or inner_product")
|
|
28
|
+
if not str(collection).strip():
|
|
29
|
+
raise ValueError("Qdrant collection must be non-empty")
|
|
30
|
+
self.client, self.collection, self.dimension = client, str(collection), dimension_value
|
|
31
|
+
self.documents, self.metric, self.vector_name = documents or {}, metric, vector_name
|
|
32
|
+
# Qdrant accepts integer IDs or UUIDs, while document stores commonly
|
|
33
|
+
# use IDs such as ``doc-17``. Keep a deterministic mapping for rows we
|
|
34
|
+
# write so arbitrary application IDs round-trip through search/retrieve.
|
|
35
|
+
self._id_map: dict[str, str] = {}
|
|
36
|
+
self._closed = False
|
|
37
|
+
|
|
38
|
+
@staticmethod
|
|
39
|
+
def _storage_id(document_id: str) -> str:
|
|
40
|
+
return str(uuid.uuid5(uuid.NAMESPACE_URL, f"embedflow:{document_id}"))
|
|
41
|
+
|
|
42
|
+
@classmethod
|
|
43
|
+
def connect(cls, path_or_url: str, collection: str, dimension: int,
|
|
44
|
+
documents: dict[str, Any] | None = None, metric: str = "cosine",
|
|
45
|
+
api_key_env: str | None = "QDRANT_API_KEY", vector_name: str | None = None) -> QdrantIndex:
|
|
46
|
+
try:
|
|
47
|
+
dimension_value = int(dimension)
|
|
48
|
+
dimension_exact = float(dimension) == dimension_value
|
|
49
|
+
except (TypeError, ValueError, OverflowError) as exc:
|
|
50
|
+
raise ValueError("Qdrant vector dimension must be a non-negative integer") from exc
|
|
51
|
+
if isinstance(dimension, bool) or not dimension_exact or dimension_value < 0:
|
|
52
|
+
raise ValueError("Qdrant vector dimension must be a non-negative integer")
|
|
53
|
+
try:
|
|
54
|
+
from qdrant_client import QdrantClient
|
|
55
|
+
except ImportError as exc:
|
|
56
|
+
raise RuntimeError("Qdrant backend requires qdrant-client; install it with `pip install qdrant-client`") from exc
|
|
57
|
+
path_or_url = str(path_or_url)
|
|
58
|
+
if "://" not in path_or_url and not path_or_url.startswith("http"):
|
|
59
|
+
client = QdrantClient(path=path_or_url)
|
|
60
|
+
else:
|
|
61
|
+
api_key = os.environ.get(api_key_env) if api_key_env else None
|
|
62
|
+
client = QdrantClient(url=path_or_url, api_key=api_key)
|
|
63
|
+
if dimension_value == 0:
|
|
64
|
+
try:
|
|
65
|
+
info = client.get_collection(collection)
|
|
66
|
+
vectors_config = getattr(getattr(info, "config", None), "params", None)
|
|
67
|
+
vectors_config = getattr(vectors_config, "vectors", None)
|
|
68
|
+
if isinstance(vectors_config, dict):
|
|
69
|
+
vectors_config = vectors_config.get(vector_name) if vector_name else next(iter(vectors_config.values()), None)
|
|
70
|
+
dimension = int(getattr(vectors_config, "size", 0) or 0)
|
|
71
|
+
if dimension <= 0:
|
|
72
|
+
raise ValueError("could not infer Qdrant collection vector dimension; set source.dimension")
|
|
73
|
+
except Exception:
|
|
74
|
+
close = getattr(client, "close", None)
|
|
75
|
+
if callable(close):
|
|
76
|
+
close()
|
|
77
|
+
raise
|
|
78
|
+
return cls(client, collection, dimension, documents, metric, vector_name)
|
|
79
|
+
|
|
80
|
+
@classmethod
|
|
81
|
+
def build(cls, path_or_url: str, collection: str, vectors: np.ndarray, ids: list[str],
|
|
82
|
+
documents: dict[str, Any] | None = None, metric: str = "cosine",
|
|
83
|
+
api_key_env: str | None = "QDRANT_API_KEY", vector_name: str | None = None) -> QdrantIndex:
|
|
84
|
+
values = np.asarray(vectors, dtype="float32")
|
|
85
|
+
metric = str(metric).lower()
|
|
86
|
+
if metric not in {"cosine", "dot", "inner_product"}:
|
|
87
|
+
raise ValueError("metric must be cosine, dot, or inner_product")
|
|
88
|
+
if values.ndim != 2 or values.shape[0] == 0 or values.shape[1] < 1 or values.shape[0] != len(ids):
|
|
89
|
+
raise ValueError("vectors must be a non-empty 2-D matrix aligned with IDs")
|
|
90
|
+
if len(set(map(str, ids))) != len(ids):
|
|
91
|
+
raise ValueError("duplicate document IDs are not allowed")
|
|
92
|
+
if not np.isfinite(values).all():
|
|
93
|
+
raise ValueError("vectors must be finite")
|
|
94
|
+
if metric == "cosine" and np.any(np.linalg.norm(values, axis=1) <= 1e-12):
|
|
95
|
+
raise ValueError("cosine vectors must be non-zero")
|
|
96
|
+
result = cls.connect(path_or_url, collection, int(values.shape[1]), documents, metric,
|
|
97
|
+
api_key_env=api_key_env, vector_name=vector_name)
|
|
98
|
+
try:
|
|
99
|
+
from qdrant_client.models import Distance, VectorParams
|
|
100
|
+
distance = Distance.COSINE if metric == "cosine" else Distance.DOT
|
|
101
|
+
vectors_config = (VectorParams(size=result.dimension, distance=distance) if not vector_name
|
|
102
|
+
else {vector_name: VectorParams(size=result.dimension, distance=distance)})
|
|
103
|
+
if hasattr(result.client, "collection_exists"):
|
|
104
|
+
if result.client.collection_exists(collection):
|
|
105
|
+
result.client.delete_collection(collection)
|
|
106
|
+
result.client.create_collection(collection_name=collection, vectors_config=vectors_config)
|
|
107
|
+
else: # qdrant-client versions before collection_exists
|
|
108
|
+
result.client.recreate_collection(collection_name=collection, vectors_config=vectors_config)
|
|
109
|
+
payloads = [{"text": documents[str(i)]} for i in ids] if documents else None
|
|
110
|
+
result.upsert(ids, values, payloads)
|
|
111
|
+
except ImportError as exc:
|
|
112
|
+
result.close()
|
|
113
|
+
raise RuntimeError("qdrant-client is required") from exc
|
|
114
|
+
except Exception:
|
|
115
|
+
result.close()
|
|
116
|
+
raise
|
|
117
|
+
return result
|
|
118
|
+
|
|
119
|
+
def search(self, query_vector: np.ndarray, k: int) -> list[SearchHit]:
|
|
120
|
+
if self._closed:
|
|
121
|
+
raise RuntimeError("Qdrant index is closed")
|
|
122
|
+
k = validate_k(k)
|
|
123
|
+
value = validate_query_vector(query_vector, self.dimension).tolist()
|
|
124
|
+
# qdrant-client renamed search -> query_points; support both APIs.
|
|
125
|
+
if hasattr(self.client, "query_points"):
|
|
126
|
+
kwargs = {"collection_name": self.collection, "query": value, "limit": k}
|
|
127
|
+
if self.vector_name:
|
|
128
|
+
kwargs["using"] = self.vector_name
|
|
129
|
+
result = self.client.query_points(**kwargs).points
|
|
130
|
+
else:
|
|
131
|
+
kwargs = {"collection_name": self.collection, "query_vector": value, "limit": k}
|
|
132
|
+
if self.vector_name:
|
|
133
|
+
try:
|
|
134
|
+
result = self.client.search(**kwargs, using=self.vector_name)
|
|
135
|
+
except TypeError:
|
|
136
|
+
result = self.client.search(collection_name=self.collection,
|
|
137
|
+
query_vector=(self.vector_name, value), limit=k)
|
|
138
|
+
else:
|
|
139
|
+
result = self.client.search(**kwargs)
|
|
140
|
+
hits = []
|
|
141
|
+
for rank, p in enumerate(result):
|
|
142
|
+
payload = p.payload or {}
|
|
143
|
+
document_id = str(payload.get("_embedflow_document_id", p.id))
|
|
144
|
+
self._id_map[document_id] = str(p.id)
|
|
145
|
+
hits.append(SearchHit(document_id, float(p.score), rank))
|
|
146
|
+
return hits
|
|
147
|
+
|
|
148
|
+
def fetch_documents(self, ids: list[str]) -> dict[str, Any]:
|
|
149
|
+
if self._closed:
|
|
150
|
+
raise RuntimeError("Qdrant index is closed")
|
|
151
|
+
if self.documents:
|
|
152
|
+
return {str(i): self.documents[str(i)] for i in ids if str(i) in self.documents}
|
|
153
|
+
if not ids:
|
|
154
|
+
return {}
|
|
155
|
+
storage_ids = [self._id_map.get(str(x), self._storage_id(str(x))) for x in ids]
|
|
156
|
+
points = self.client.retrieve(collection_name=self.collection, ids=storage_ids, with_payload=True)
|
|
157
|
+
output = {}
|
|
158
|
+
for p in points:
|
|
159
|
+
payload = p.payload or {}
|
|
160
|
+
document_id = str(payload.get("_embedflow_document_id", p.id))
|
|
161
|
+
output[document_id] = payload.get("text", payload)
|
|
162
|
+
return output
|
|
163
|
+
|
|
164
|
+
def size(self) -> int:
|
|
165
|
+
if self._closed:
|
|
166
|
+
raise RuntimeError("Qdrant index is closed")
|
|
167
|
+
info = self.client.get_collection(self.collection)
|
|
168
|
+
return int(getattr(info, "points_count", 0) or 0)
|
|
169
|
+
|
|
170
|
+
def metadata(self) -> dict[str, Any]:
|
|
171
|
+
if self._closed:
|
|
172
|
+
raise RuntimeError("Qdrant index is closed")
|
|
173
|
+
return {"backend": "qdrant", "collection": self.collection, "dimension": self.dimension,
|
|
174
|
+
"metric": self.metric, "vector_name": self.vector_name, "size": self.size()}
|
|
175
|
+
|
|
176
|
+
def close(self) -> None:
|
|
177
|
+
"""Release a local file lock or remote client connection.
|
|
178
|
+
|
|
179
|
+
``QdrantClient.close`` is available in current qdrant-client releases;
|
|
180
|
+
older clients simply have no close method, so teardown remains safe.
|
|
181
|
+
The method is idempotent to support engine and application shutdown
|
|
182
|
+
paths both calling it defensively.
|
|
183
|
+
"""
|
|
184
|
+
if self._closed:
|
|
185
|
+
return
|
|
186
|
+
self._closed = True
|
|
187
|
+
close = getattr(self.client, "close", None)
|
|
188
|
+
if callable(close):
|
|
189
|
+
close()
|
|
190
|
+
|
|
191
|
+
def upsert(self, ids: list[str], vectors: np.ndarray, payloads: list[dict[str, Any]] | None = None) -> None:
|
|
192
|
+
try:
|
|
193
|
+
from qdrant_client.models import PointStruct
|
|
194
|
+
except ImportError as exc:
|
|
195
|
+
raise RuntimeError("qdrant-client is required") from exc
|
|
196
|
+
ids = [str(value) for value in ids]
|
|
197
|
+
values = np.asarray(vectors, dtype="float32")
|
|
198
|
+
if values.ndim != 2 or values.shape != (len(ids), self.dimension):
|
|
199
|
+
raise ValueError(f"Qdrant vectors must have shape ({len(ids)}, {self.dimension})")
|
|
200
|
+
if not ids or len(set(ids)) != len(ids):
|
|
201
|
+
raise ValueError("Qdrant upsert requires non-empty unique document IDs")
|
|
202
|
+
if not np.isfinite(values).all():
|
|
203
|
+
raise ValueError("Qdrant vectors must be finite")
|
|
204
|
+
if self.metric == "cosine" and np.any(np.linalg.norm(values, axis=1) <= 1e-12):
|
|
205
|
+
raise ValueError("Qdrant cosine vectors must be non-zero")
|
|
206
|
+
if payloads is not None and len(payloads) != len(ids):
|
|
207
|
+
raise ValueError("Qdrant payload count must match document IDs")
|
|
208
|
+
payloads = payloads or [{} for _ in ids]
|
|
209
|
+
storage_ids = []
|
|
210
|
+
normalized_payloads = []
|
|
211
|
+
for document_id, payload in zip(ids, payloads):
|
|
212
|
+
original_id = str(document_id)
|
|
213
|
+
storage_id = self._storage_id(original_id)
|
|
214
|
+
self._id_map[original_id] = storage_id
|
|
215
|
+
storage_ids.append(storage_id)
|
|
216
|
+
if payload is not None and not isinstance(payload, dict):
|
|
217
|
+
raise ValueError("Qdrant payloads must be objects")
|
|
218
|
+
item = dict(payload or {})
|
|
219
|
+
item.setdefault("_embedflow_document_id", original_id)
|
|
220
|
+
normalized_payloads.append(item)
|
|
221
|
+
self.client.upsert(collection_name=self.collection,
|
|
222
|
+
points=[PointStruct(id=storage_id,
|
|
223
|
+
vector=({self.vector_name: v.tolist()} if self.vector_name else v.tolist()),
|
|
224
|
+
payload=payload)
|
|
225
|
+
for storage_id, v, payload in zip(storage_ids, values, normalized_payloads)])
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from collections import defaultdict
|
|
4
|
+
from collections.abc import Iterable
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
STAGES = ("source_query_encode_ms", "source_ann_ms", "target_query_encode_ms", "cache_lookup_ms",
|
|
10
|
+
"synchronous_target_encode_ms", "target_score_ms", "topk_ms", "total_ms")
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def summarize(values: Iterable[float]) -> dict[str, float | int]:
|
|
14
|
+
arr = np.asarray(list(values), dtype="float64")
|
|
15
|
+
if arr.size == 0: return {"count": 0}
|
|
16
|
+
if not np.isfinite(arr).all() or (arr < 0).any(): raise ValueError("latency values must be finite and non-negative")
|
|
17
|
+
return {"count": int(arr.size), "mean_ms": float(np.mean(arr)), "p50_ms": float(np.quantile(arr, .50)),
|
|
18
|
+
"p90_ms": float(np.quantile(arr, .90)), "p95_ms": float(np.quantile(arr, .95)),
|
|
19
|
+
"p99_ms": float(np.quantile(arr, .99)), "min_ms": float(np.min(arr)), "max_ms": float(np.max(arr)),
|
|
20
|
+
"std_ms": float(np.std(arr, ddof=1)) if arr.size > 1 else 0.0,
|
|
21
|
+
# A zero-duration synthetic fixture has no meaningful finite QPS;
|
|
22
|
+
# return 0 rather than leaking infinity into JSON reports.
|
|
23
|
+
"qps": float(1000.0 / np.mean(arr)) if float(np.mean(arr)) > 0 else 0.0}
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def aggregate_records(records: list[dict[str, Any]], group_keys: tuple[str, ...] = ("mode", "K", "nprobe")) -> list[dict[str, Any]]:
|
|
27
|
+
grouped: dict[tuple[Any, ...], list[dict[str, Any]]] = defaultdict(list)
|
|
28
|
+
malformed = 0
|
|
29
|
+
for row in records:
|
|
30
|
+
if not isinstance(row, dict):
|
|
31
|
+
malformed += 1
|
|
32
|
+
continue
|
|
33
|
+
try:
|
|
34
|
+
values = [float(row.get(stage, 0.0)) for stage in STAGES]
|
|
35
|
+
if not np.isfinite(values).all() or (np.asarray(values) < 0).any():
|
|
36
|
+
raise ValueError
|
|
37
|
+
except (TypeError, ValueError):
|
|
38
|
+
# Telemetry is append-only and may contain a partially written or
|
|
39
|
+
# hand-edited line. Skip that row while making the omission visible.
|
|
40
|
+
malformed += 1
|
|
41
|
+
continue
|
|
42
|
+
grouped[tuple(row.get(k) for k in group_keys)].append(row)
|
|
43
|
+
out = []
|
|
44
|
+
for key, rows in sorted(grouped.items(), key=lambda x: tuple(str(v) for v in x[0])):
|
|
45
|
+
result = dict(zip(group_keys, key))
|
|
46
|
+
for stage in STAGES: result.update({f"{stage}_{metric}": value for metric, value in summarize(float(r.get(stage, 0.0)) for r in rows).items()})
|
|
47
|
+
if malformed:
|
|
48
|
+
result["malformed_rows_skipped"] = malformed
|
|
49
|
+
out.append(result)
|
|
50
|
+
return out
|
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import random
|
|
5
|
+
from collections.abc import Iterable, Mapping
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
K_VALUES = (10, 20, 50, 100, 200, 500)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _research_features(records: list[tuple[str, float, int]], target_scores: dict[str, float], ks: tuple[int, ...]) -> dict[str, float]:
|
|
15
|
+
"""Delegate the frozen feature implementation when this checkout has it."""
|
|
16
|
+
try:
|
|
17
|
+
if max(ks) < 500:
|
|
18
|
+
raise ImportError("the frozen helper requires the complete K<=500 feature set")
|
|
19
|
+
from src.probe_features import build_features
|
|
20
|
+
return build_features(records, target_scores, ks=ks)
|
|
21
|
+
except (ImportError, ModuleNotFoundError):
|
|
22
|
+
final_ids = [x[0] for x in _rerank(records, target_scores, max(ks))[:10]]
|
|
23
|
+
top50 = [x[0] for x in _rerank(records, target_scores, min(50, len(records)))[:10]]
|
|
24
|
+
top500 = [x[0] for x in _rerank(records, target_scores, min(500, len(records)))[:10]]
|
|
25
|
+
stability = len(set(top50) & set(final_ids)) / max(1, len(final_ids))
|
|
26
|
+
return {"probe_residual_tail_50_mean": 1 - stability, "stability_to_500_50_mean": stability,
|
|
27
|
+
"deepest_p90": float(max((x[2] for x in _rerank(records, target_scores, min(500, len(records)))[:10]), default=0)),
|
|
28
|
+
"late_tail_area": float(len(set(top500) - set(top50)) / 10), "last_shell_any_rate": float(bool(set(top500) - set(top50))),
|
|
29
|
+
"fraction_margin_nonpositive": 0.0}
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _rerank(records: list[tuple[str, float, int]], scores: dict[str, float], k: int) -> list[tuple[str, float, int]]:
|
|
33
|
+
subset = records[:min(k, len(records))]
|
|
34
|
+
return sorted(subset, key=lambda x: (-float(scores[x[0]]), int(x[2])))
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def _mean_features(rows: list[dict[str, float]]) -> dict[str, float]:
|
|
38
|
+
if not rows:
|
|
39
|
+
raise ValueError("cannot aggregate empty feature rows")
|
|
40
|
+
keys = sorted({k for r in rows for k in r})
|
|
41
|
+
output: dict[str, float] = {}
|
|
42
|
+
for key in keys:
|
|
43
|
+
values = [float(r.get(key, 0.0)) for r in rows]
|
|
44
|
+
if not np.isfinite(values).all():
|
|
45
|
+
raise ValueError(f"feature {key!r} contains non-finite values")
|
|
46
|
+
output[key] = float(np.mean(values))
|
|
47
|
+
return output
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def recommend_k(per_query: list[dict[str, Any]], kmax: int) -> tuple[int, str]:
|
|
51
|
+
if isinstance(kmax, bool) or int(kmax) != kmax or int(kmax) < 10:
|
|
52
|
+
raise ValueError("kmax must be an integer >= 10")
|
|
53
|
+
if not per_query: return min(50, int(kmax)), "No probe rows; using the configured conservative default."
|
|
54
|
+
for k in K_VALUES:
|
|
55
|
+
if k > kmax: continue
|
|
56
|
+
try:
|
|
57
|
+
vals = [float(row.get(f"stability_{k}", row.get("stability_to_500_50_mean", 0.0))) for row in per_query]
|
|
58
|
+
except (TypeError, ValueError) as exc:
|
|
59
|
+
raise ValueError(f"probe stability at K={k} must be numeric") from exc
|
|
60
|
+
if not np.isfinite(vals).all() or any(value < 0 or value > 1 for value in vals):
|
|
61
|
+
raise ValueError(f"probe stability at K={k} must be finite and in [0, 1]")
|
|
62
|
+
if vals and float(np.mean(vals)) >= 0.90:
|
|
63
|
+
return k, f"Mean finite-pool top-10 stability at K={k} is {float(np.mean(vals)):.3f}."
|
|
64
|
+
return min(200, int(kmax)), "No tested K reached the 0.90 finite-pool stability heuristic; review before deployment."
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def run_probe(source_model: Any, target_model: Any, index: Any, documents: Any,
|
|
68
|
+
queries: Iterable[tuple[str, str]], kmax: int = 500, seed: int = 42,
|
|
69
|
+
limit: int | None = None,
|
|
70
|
+
target_document_vectors: Mapping[str, np.ndarray] | None = None) -> dict[str, Any]:
|
|
71
|
+
"""Run a finite candidate compatibility probe with frozen T2-v1 decision logic."""
|
|
72
|
+
if isinstance(kmax, bool) or int(kmax) != kmax or int(kmax) < 10:
|
|
73
|
+
raise ValueError("probe kmax must be an integer >= 10")
|
|
74
|
+
if limit is not None and (isinstance(limit, bool) or int(limit) != limit or int(limit) < 1):
|
|
75
|
+
raise ValueError("probe limit must be a positive integer when provided")
|
|
76
|
+
rows = list(queries)
|
|
77
|
+
normalized_rows: list[tuple[str, str]] = []
|
|
78
|
+
seen_query_ids: set[str] = set()
|
|
79
|
+
for item in rows:
|
|
80
|
+
if not isinstance(item, (tuple, list)) or len(item) != 2:
|
|
81
|
+
raise ValueError("probe queries must be (query_id, text) pairs")
|
|
82
|
+
query_id, text = str(item[0]), item[1]
|
|
83
|
+
if not isinstance(text, str) or not text.strip():
|
|
84
|
+
raise ValueError(f"probe query {query_id!r} must contain non-empty text")
|
|
85
|
+
if query_id in seen_query_ids:
|
|
86
|
+
raise ValueError(f"duplicate probe query ID {query_id!r}")
|
|
87
|
+
seen_query_ids.add(query_id)
|
|
88
|
+
normalized_rows.append((query_id, text))
|
|
89
|
+
rows = normalized_rows
|
|
90
|
+
if not rows:
|
|
91
|
+
raise ValueError("probe queries must contain at least one query")
|
|
92
|
+
rng = random.Random(seed); rng.shuffle(rows)
|
|
93
|
+
rows = rows[:limit] if limit else rows
|
|
94
|
+
kmax = min(int(kmax), index.size())
|
|
95
|
+
if kmax < 10: raise ValueError("compatibility probe needs at least 10 index candidates")
|
|
96
|
+
try:
|
|
97
|
+
from src.t2_v1 import verify_t2_hash
|
|
98
|
+
root = Path(__file__).resolve().parents[2]
|
|
99
|
+
package_root = Path(__file__).resolve().parents[1]
|
|
100
|
+
if (root / "frozen").exists():
|
|
101
|
+
verify_t2_hash(root)
|
|
102
|
+
elif (package_root / "frozen").exists():
|
|
103
|
+
verify_t2_hash(package_root)
|
|
104
|
+
except (ImportError, FileNotFoundError):
|
|
105
|
+
pass
|
|
106
|
+
per_query = []
|
|
107
|
+
for query_id, text in rows:
|
|
108
|
+
sq = np.asarray(source_model.encode_queries([text])[0], dtype="float32")
|
|
109
|
+
if sq.ndim != 1 or sq.shape[0] != int(source_model.dimension) or not np.isfinite(sq).all():
|
|
110
|
+
raise ValueError("source query encoder returned an invalid vector")
|
|
111
|
+
candidates = index.search(sq, kmax)
|
|
112
|
+
ids = [x.document_id for x in candidates]
|
|
113
|
+
if not ids:
|
|
114
|
+
raise ValueError("source index returned no candidates")
|
|
115
|
+
if target_document_vectors is None:
|
|
116
|
+
docs = documents.get(ids)
|
|
117
|
+
tv = np.asarray(target_model.encode_documents([docs[x] for x in ids]), dtype="float32")
|
|
118
|
+
else:
|
|
119
|
+
missing_vectors = [document_id for document_id in ids if document_id not in target_document_vectors]
|
|
120
|
+
if missing_vectors:
|
|
121
|
+
raise ValueError(f"target probe vectors are missing {len(missing_vectors)} candidate IDs (e.g. {missing_vectors[:3]})")
|
|
122
|
+
tv = np.asarray([target_document_vectors[document_id] for document_id in ids], dtype="float32")
|
|
123
|
+
tq = np.asarray(target_model.encode_queries([text])[0], dtype="float32")
|
|
124
|
+
if tv.ndim != 2 or tv.shape != (len(ids), int(target_model.dimension)) or not np.isfinite(tv).all():
|
|
125
|
+
raise ValueError("target document dimension or finiteness mismatch in probe")
|
|
126
|
+
if tq.ndim != 1 or tq.shape[0] != int(target_model.dimension) or not np.isfinite(tq).all():
|
|
127
|
+
raise ValueError("target query encoder returned an invalid vector")
|
|
128
|
+
scores = {did: float(vec @ tq) for did, vec in zip(ids, tv)}
|
|
129
|
+
features = _research_features([(x.document_id, x.score, x.source_rank) for x in candidates], scores, tuple(k for k in K_VALUES if k <= kmax))
|
|
130
|
+
features["query_id"] = str(query_id)
|
|
131
|
+
features["target_scores"] = scores
|
|
132
|
+
# Preserve per-K finite-pool stability for K recommendation.
|
|
133
|
+
final = [x[0] for x in _rerank([(x.document_id, x.score, x.source_rank) for x in candidates], scores, kmax)[:10]]
|
|
134
|
+
for k in K_VALUES:
|
|
135
|
+
if k <= kmax:
|
|
136
|
+
got = [x[0] for x in _rerank([(x.document_id, x.score, x.source_rank) for x in candidates], scores, k)[:10]]
|
|
137
|
+
features[f"stability_{k}"] = len(set(got) & set(final)) / max(1, len(final))
|
|
138
|
+
per_query.append(features)
|
|
139
|
+
scalar_rows = [{k: v for k, v in row.items() if isinstance(v, (int, float, np.number))} for row in per_query]
|
|
140
|
+
aggregate = _mean_features(scalar_rows)
|
|
141
|
+
try:
|
|
142
|
+
from src.t2_v1 import decide
|
|
143
|
+
diagnostic = decide(aggregate)
|
|
144
|
+
implementation = "src.t2_v1.decide"
|
|
145
|
+
except (ImportError, KeyError) as exc:
|
|
146
|
+
raise RuntimeError("frozen T2-v1 implementation is unavailable; refusing to invent a rule") from exc
|
|
147
|
+
recommended, rationale = recommend_k(per_query, kmax)
|
|
148
|
+
return {"diagnostic": diagnostic, "recommended_k": recommended, "rationale": rationale,
|
|
149
|
+
"features": aggregate, "per_query": per_query, "queries": len(per_query),
|
|
150
|
+
"kmax": kmax, "seed": seed, "implementation": implementation,
|
|
151
|
+
"warning": "T2-v1 is an empirical finite-tail diagnostic, not a mathematical guarantee."}
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def save_probe(result: dict[str, Any], path: str | Path) -> None:
|
|
155
|
+
path = Path(path); path.parent.mkdir(parents=True, exist_ok=True)
|
|
156
|
+
path.write_text(json.dumps(result, indent=2, ensure_ascii=False, default=float))
|