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.
Files changed (64) hide show
  1. embedflow/__init__.py +25 -0
  2. embedflow/__main__.py +3 -0
  3. embedflow/analysis.py +192 -0
  4. embedflow/cache/__init__.py +4 -0
  5. embedflow/cache/base.py +28 -0
  6. embedflow/cache/persistent_cache.py +198 -0
  7. embedflow/cli.py +1200 -0
  8. embedflow/compatibility/__init__.py +28 -0
  9. embedflow/compatibility/candidate_gap.py +105 -0
  10. embedflow/compatibility/containment.py +17 -0
  11. embedflow/compatibility/evaluate.py +319 -0
  12. embedflow/compatibility/metrics.py +75 -0
  13. embedflow/compatibility/migration_depth.py +67 -0
  14. embedflow/compatibility/probe.py +34 -0
  15. embedflow/compatibility/report.py +102 -0
  16. embedflow/compatibility/t2.py +64 -0
  17. embedflow/config.py +455 -0
  18. embedflow/data/__init__.py +1 -0
  19. embedflow/data/registry/__init__.py +1 -0
  20. embedflow/data/registry/benchmark_profiles.jsonl +3 -0
  21. embedflow/data/registry/checksums.sha256 +4 -0
  22. embedflow/data/registry/migrations.jsonl +15 -0
  23. embedflow/data/registry/registry_manifest.json +16 -0
  24. embedflow/data/registry/research_summaries.json +55 -0
  25. embedflow/data/registry/schema_version.json +5 -0
  26. embedflow/frozen/T2_V1_FROZEN_SPEC.md +71 -0
  27. embedflow/frozen/T2_V1_FROZEN_SPEC.sha256 +1 -0
  28. embedflow/indexes/__init__.py +5 -0
  29. embedflow/indexes/base.py +60 -0
  30. embedflow/indexes/faiss_backend.py +240 -0
  31. embedflow/indexes/qdrant_backend.py +225 -0
  32. embedflow/metrics/__init__.py +3 -0
  33. embedflow/metrics/latency.py +50 -0
  34. embedflow/migration/__init__.py +3 -0
  35. embedflow/migration/compatibility.py +156 -0
  36. embedflow/migration/facade.py +312 -0
  37. embedflow/migration/materializer.py +190 -0
  38. embedflow/migration/planner.py +78 -0
  39. embedflow/migration/state.py +81 -0
  40. embedflow/models/__init__.py +4 -0
  41. embedflow/models/base.py +31 -0
  42. embedflow/models/huggingface.py +226 -0
  43. embedflow/registry/__init__.py +47 -0
  44. embedflow/registry/loader.py +785 -0
  45. embedflow/registry/matcher.py +197 -0
  46. embedflow/registry/schema.py +266 -0
  47. embedflow/runtime.py +115 -0
  48. embedflow/serving/__init__.py +3 -0
  49. embedflow/serving/api.py +161 -0
  50. embedflow/serving/engine.py +222 -0
  51. embedflow/serving/factory.py +3 -0
  52. embedflow/serving/schemas.py +39 -0
  53. embedflow-0.1.0.dist-info/METADATA +210 -0
  54. embedflow-0.1.0.dist-info/RECORD +64 -0
  55. embedflow-0.1.0.dist-info/WHEEL +5 -0
  56. embedflow-0.1.0.dist-info/entry_points.txt +2 -0
  57. embedflow-0.1.0.dist-info/licenses/LICENSE +178 -0
  58. embedflow-0.1.0.dist-info/top_level.txt +2 -0
  59. src/__init__.py +1 -0
  60. src/embed.py +123 -0
  61. src/probe_features.py +24 -0
  62. src/storage.py +51 -0
  63. src/t2_v1.py +21 -0
  64. 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,3 @@
1
+ from .latency import STAGES, aggregate_records, summarize
2
+
3
+ __all__ = ["STAGES", "aggregate_records", "summarize"]
@@ -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,3 @@
1
+ from .state import DocumentStore, MigrationState
2
+
3
+ __all__ = ["DocumentStore", "MigrationState"]
@@ -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))