From ea78bf662e225d8008c12993e7cbb36457f39264 Mon Sep 17 00:00:00 2001 From: triplec <202340740213@stu.hunau.edu.cn> Date: Fri, 17 Jul 2026 04:11:15 -0700 Subject: [PATCH 1/2] feat(retrieval): add pluggable cross-encoder reranker --- .env.example | 5 +- docs/en/concepts/reranker-benchmark.md | 28 ++++ docs/en/concepts/retrieval-model.md | 22 ++- docs/en/reference/settings.md | 5 +- docs/zh/reference/settings.md | 5 +- pyproject.toml | 2 + scripts/benchmark_rerankers.py | 134 +++++++++++++++++ src/contextseek/client/contextseek.py | 8 +- src/contextseek/config/__init__.py | 8 +- src/contextseek/config/factory.py | 25 ++++ src/contextseek/config/settings.py | 3 + src/contextseek/config/strategies.py | 2 +- src/contextseek/retrieval/__init__.py | 2 + src/contextseek/retrieval/components.py | 71 +++++++++ .../retrieval/test_reranker_features.py | 62 +++++++- tests/unit_tests/test_settings.py | 41 +++++ uv.lock | 141 ++++++++---------- 17 files changed, 480 insertions(+), 84 deletions(-) create mode 100644 docs/en/concepts/reranker-benchmark.md create mode 100644 scripts/benchmark_rerankers.py diff --git a/.env.example b/.env.example index 48ed83d..1202eb0 100644 --- a/.env.example +++ b/.env.example @@ -134,8 +134,11 @@ RETRIEVAL_DEFAULT_K=20 # RETRIEVAL_DEFAULT_BUDGET=20 # RETRIEVAL_RECALL_ROUTES=["phrase","terms"] # RETRIEVAL_RECALL_ROUTES=["phrase","terms","vector"] -# RETRIEVAL_RERANKER_MODE=heuristic # heuristic | llm +# RETRIEVAL_RERANKER_MODE=heuristic # heuristic | llm | cross_encoder # RETRIEVAL_LLM_RERANK_TOP_N=20 +# RETRIEVAL_CROSS_ENCODER_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2 +# RETRIEVAL_CROSS_ENCODER_DEVICE= # empty = auto, or cpu/cuda +# RETRIEVAL_CROSS_ENCODER_TOP_N=20 # RETRIEVAL_VECTOR_WEIGHT=0.7 # RETRIEVAL_FTS_WEIGHT=0.3 # RETRIEVAL_MAX_CONTENT_CHARS=1200 diff --git a/docs/en/concepts/reranker-benchmark.md b/docs/en/concepts/reranker-benchmark.md new file mode 100644 index 0000000..00289e8 --- /dev/null +++ b/docs/en/concepts/reranker-benchmark.md @@ -0,0 +1,28 @@ +# Reranker benchmark + +This smoke benchmark compares the built-in heuristic and cross-encoder paths on +five small support/operations queries. Each query has five recalled candidates +and one labeled relevant candidate. It measures ranking recall and reranking +latency only; model download and startup are excluded. + +Run it with: + +```bash +uv run --extra rerank python scripts/benchmark_rerankers.py --device cpu --rounds 5 +``` + +The fixture is intentionally small and deterministic, so it is useful for +regression checks rather than as a general model-quality claim. Replace the +cases in the script with domain-specific labeled candidates before selecting a +production model. + +## Reference result + +Results will vary by CPU, device, model cache, and candidate length. The table +below records a five-round run using Python 3.13 on a Windows x86-64 CPU with +`cross-encoder/ms-marco-MiniLM-L-6-v2` after one warm-up batch. + +| Reranker | Recall@1 | Recall@3 | Median latency/query | +|----------|----------|----------|----------------------| +| Heuristic | 0.000 | 0.600 | 0.08 ms | +| Cross-encoder | 0.200 | 1.000 | 15.83 ms | diff --git a/docs/en/concepts/retrieval-model.md b/docs/en/concepts/retrieval-model.md index 795a367..2bc3d1b 100644 --- a/docs/en/concepts/retrieval-model.md +++ b/docs/en/concepts/retrieval-model.md @@ -53,7 +53,7 @@ All active routes run in parallel; their candidate sets are merged before rerank ## Reranking -After recall, candidates are scored and ranked. Two modes: +After recall, candidates are scored and ranked. Three modes: ### Heuristic reranker (default) @@ -79,6 +79,26 @@ RETRIEVAL_RERANKER_MODE=llm RETRIEVAL_LLM_RERANK_TOP_N=20 ``` +### Cross-encoder reranker + +Install the optional model dependency and select the local cross-encoder path: + +```bash +pip install "contextseek[rerank]" +``` + +```env +RETRIEVAL_RERANKER_MODE=cross_encoder +RETRIEVAL_CROSS_ENCODER_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2 +RETRIEVAL_CROSS_ENCODER_TOP_N=20 +``` + +The model scores all selected query/document pairs in one batch. Loading is +lazy, so heuristic and LLM configurations do not import Sentence Transformers. +Set `RETRIEVAL_CROSS_ENCODER_DEVICE=cpu` or `cuda` to override automatic device +selection. See the [reranker benchmark](reranker-benchmark.md) for a reproducible +quality/latency comparison. + --- ## Layer selection: summary vs. full diff --git a/docs/en/reference/settings.md b/docs/en/reference/settings.md index aa65752..67eabfb 100644 --- a/docs/en/reference/settings.md +++ b/docs/en/reference/settings.md @@ -86,8 +86,11 @@ When `SUMMARIZER_PROVIDER=llm` but no LLM is configured, the summarizer is skipp | `RETRIEVAL_LINK_BOOST` | `0.10` | Score bonus for items with supporting links | | `RETRIEVAL_LINK_REFUTE_PENALTY` | `0.40` | Score penalty for items with refuting links | | `RETRIEVAL_LINK_SUPERSEDE_PENALTY` | `0.35` | Score penalty for superseded items | -| `RETRIEVAL_RERANKER_MODE` | `heuristic` | `heuristic` or `llm` | +| `RETRIEVAL_RERANKER_MODE` | `heuristic` | `heuristic`, `llm`, or `cross_encoder` | | `RETRIEVAL_LLM_RERANK_TOP_N` | `20` | Candidate count passed to LLM reranker | +| `RETRIEVAL_CROSS_ENCODER_MODEL` | `cross-encoder/ms-marco-MiniLM-L-6-v2` | Sentence Transformers cross-encoder model | +| `RETRIEVAL_CROSS_ENCODER_DEVICE` | _(auto)_ | Model device, for example `cpu` or `cuda` | +| `RETRIEVAL_CROSS_ENCODER_TOP_N` | `20` | Candidate count passed to the cross-encoder | ## Evolution (`EVOLUTION_*`) diff --git a/docs/zh/reference/settings.md b/docs/zh/reference/settings.md index caa6101..4dd2e38 100644 --- a/docs/zh/reference/settings.md +++ b/docs/zh/reference/settings.md @@ -86,8 +86,11 @@ Provider 的 API Key(`OPENAI_API_KEY`、`DASHSCOPE_API_KEY` 等)由 LangChai | `RETRIEVAL_LINK_BOOST` | `0.10` | 有支持链接的条目的得分加成 | | `RETRIEVAL_LINK_REFUTE_PENALTY` | `0.40` | 有反驳链接的条目的得分惩罚 | | `RETRIEVAL_LINK_SUPERSEDE_PENALTY` | `0.35` | 已被替代条目的得分惩罚 | -| `RETRIEVAL_RERANKER_MODE` | `heuristic` | `heuristic` 或 `llm` | +| `RETRIEVAL_RERANKER_MODE` | `heuristic` | `heuristic`、`llm` 或 `cross_encoder` | | `RETRIEVAL_LLM_RERANK_TOP_N` | `20` | 传给 LLM 重排器的候选数量 | +| `RETRIEVAL_CROSS_ENCODER_MODEL` | `cross-encoder/ms-marco-MiniLM-L-6-v2` | Sentence Transformers cross-encoder 模型 | +| `RETRIEVAL_CROSS_ENCODER_DEVICE` | _(自动)_ | 模型设备,例如 `cpu` 或 `cuda` | +| `RETRIEVAL_CROSS_ENCODER_TOP_N` | `20` | 传给 cross-encoder 的候选数量 | ## 演化(`EVOLUTION_*`) diff --git a/pyproject.toml b/pyproject.toml index 68c3bd8..fffc6f2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -50,11 +50,13 @@ oceanbase = ["pyobvector>=0.1.0", "sqlalchemy>=2.0"] daemon = ["watchdog>=4.0"] seekdb = ["pyseekdb"] powermem = ["powermem>=1.1.1"] +rerank = ["sentence-transformers>=3.0.0"] all = [ "fastapi>=0.110.0", "uvicorn>=0.30.0", "pydantic>=2.7.0", "cryptography>=42.0.0", "langchain-core>=0.3.0", "langchain>=0.3.0", "langgraph>=0.2.0", "langchain-openai>=0.3.0", "langchain-community>=0.3.0", "dashscope>=1.20.0", + "sentence-transformers>=3.0.0", "watchdog>=4.0", "pyseekdb", ] diff --git a/scripts/benchmark_rerankers.py b/scripts/benchmark_rerankers.py new file mode 100644 index 0000000..3c1d01e --- /dev/null +++ b/scripts/benchmark_rerankers.py @@ -0,0 +1,134 @@ +"""Small reproducible recall/latency benchmark for retrieval rerankers.""" + +from __future__ import annotations + +import argparse +from dataclasses import dataclass +from statistics import median +from time import perf_counter + +from contextseek.config.strategies import RetrievalStrategy +from contextseek.retrieval.components import CrossEncoderReranker, HeuristicReranker + + +@dataclass(frozen=True) +class Case: + query: str + relevant_id: str + documents: tuple[tuple[str, str], ...] + + +CASES = ( + Case( + "How can I reset a forgotten password?", + "password", + ( + ("analytics", "Password reset analytics count forgotten-password requests."), + ("billing", "Invoices can be downloaded from account billing settings."), + ("password", "Use the forgot password link to receive a reset email."), + ("profile", "Change your display name from the profile page."), + ("session", "Active browser sessions can be reviewed by an administrator."), + ), + ), + Case( + "Why did the deployment run out of memory?", + "oom", + ( + ("policy", "The deployment memory policy was reviewed last quarter."), + ("timeout", "The deployment failed because the health check timed out."), + ("oom", "The container was killed after exceeding its memory limit."), + ("network", "The release could not resolve the package registry hostname."), + ("version", "The release changed the application version label."), + ), + ), + Case( + "Where do I rotate an API credential?", + "key", + ( + ("policy", "The API credential rotation policy is reviewed annually."), + ("key", "Create and revoke access tokens from the security console."), + ("logs", "Audit logs retain administrative actions for ninety days."), + ("quota", "Request a higher rate limit from the usage page."), + ("webhook", "Webhooks notify external systems when records change."), + ), + ), + Case( + "Can deleted records be recovered?", + "restore", + ( + ("report", "Deleted records are excluded from active-record reports."), + ("export", "Export active records as a CSV file."), + ("restore", "Restore soft-deleted items from the recycle bin for 30 days."), + ("retention", "Archived logs are retained for compliance."), + ("schema", "Custom fields can be added to record schemas."), + ), + ), + Case( + "How do I reduce slow database queries?", + "index", + ( + ("monitor", "A dashboard lists slow database queries and their duration."), + ("backup", "Schedule nightly snapshots of the database."), + ("index", "Add an index for columns used by frequent query filters."), + ("replica", "Read replicas improve availability during maintenance."), + ("access", "Database roles control which tables a user can access."), + ), + ), +) + + +def _candidates(case: Case) -> list[dict[str, object]]: + return [ + {"id": item_id, "content": content, "score": 0.5, "stage": "skill"} + for item_id, content in case.documents + ] + + +def evaluate(reranker, *, rounds: int) -> tuple[float, float, float]: + strategy = RetrievalStrategy() + top_one = 0 + top_three = 0 + latencies_ms: list[float] = [] + for _ in range(rounds): + for case in CASES: + started = perf_counter() + ranked = reranker.rerank( + _candidates(case), query=case.query, strategy=strategy + ) + latencies_ms.append((perf_counter() - started) * 1000) + ids = [str(item["id"]) for item in ranked] + top_one += case.relevant_id in ids[:1] + top_three += case.relevant_id in ids[:3] + total = len(CASES) * rounds + return top_one / total, top_three / total, median(latencies_ms) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument( + "--model", default="cross-encoder/ms-marco-MiniLM-L-6-v2" + ) + parser.add_argument("--device", default="cpu") + parser.add_argument("--rounds", type=int, default=5) + args = parser.parse_args() + + cross_encoder = CrossEncoderReranker(args.model, device=args.device) + cross_encoder.rerank( + _candidates(CASES[0]), query=CASES[0].query, strategy=RetrievalStrategy() + ) + + print("reranker\trecall@1\trecall@3\tmedian_ms") + for name, reranker in ( + ("heuristic", HeuristicReranker()), + ("cross_encoder", cross_encoder), + ): + recall_at_one, recall_at_three, latency = evaluate( + reranker, rounds=max(1, args.rounds) + ) + print( + f"{name}\t{recall_at_one:.3f}\t{recall_at_three:.3f}\t{latency:.2f}" + ) + + +if __name__ == "__main__": + main() diff --git a/src/contextseek/client/contextseek.py b/src/contextseek/client/contextseek.py index 774a988..4fba901 100644 --- a/src/contextseek/client/contextseek.py +++ b/src/contextseek/client/contextseek.py @@ -387,6 +387,9 @@ class ContextSeek: llm: Any | None = None """Optional shared LLM for advanced ranking/evolution/classification hooks.""" + reranker: Any | None = None + """Optional retrieval reranker implementing the Reranker protocol.""" + llm_prompts: LLMPromptTemplates = field(default_factory=LLMPromptTemplates) """Prompt templates used by all LLM-assisted flows.""" @@ -613,7 +616,7 @@ def retrieve( from contextseek.retrieval.orchestrator import RetrievalOrchestrator from contextseek.retrieval.components import LLMReranker - reranker = None + reranker = self.reranker if self._llm_rerank_enabled and self.llm is not None: reranker = LLMReranker( score_fn=self._score_relevance_with_llm, @@ -2414,6 +2417,7 @@ def from_settings( from contextseek.config.factory import ( build_embedder, build_llm, + build_reranker, build_summarizer, ) @@ -2490,6 +2494,7 @@ def _seekdb_embed(text: str, _ef: Any = _seekdb_ef) -> list[float]: llm=shared_llm, prompt_templates=llm_prompts, ) + reranker = build_reranker(settings.retrieval) llm_rerank_enabled = ( shared_llm is not None and settings.retrieval.reranker_mode.lower() == "llm" @@ -2562,6 +2567,7 @@ def _seekdb_embed(text: str, _ef: Any = _seekdb_ef) -> list[float]: embedder=embedder, summarizer=summarizer, llm=shared_llm, + reranker=reranker, llm_prompts=llm_prompts, evolution_engine=evolution_engine, audit_log=audit_log, diff --git a/src/contextseek/config/__init__.py b/src/contextseek/config/__init__.py index a3e10dd..8297d14 100644 --- a/src/contextseek/config/__init__.py +++ b/src/contextseek/config/__init__.py @@ -18,7 +18,12 @@ from contextseek.config.settings import nested_section_config from contextseek.config.settings import settings_config from contextseek.config.settings import to_strategy_config -from contextseek.config.factory import build_embedder, build_llm, build_summarizer +from contextseek.config.factory import ( + build_embedder, + build_llm, + build_reranker, + build_summarizer, +) __all__ = [ "EvolutionStrategy", @@ -36,6 +41,7 @@ "WriteStrategy", "build_embedder", "build_llm", + "build_reranker", "build_summarizer", "default_strategy_config", "HYBRID_RETRIEVAL_STRATEGY", diff --git a/src/contextseek/config/factory.py b/src/contextseek/config/factory.py index 3d4fe5c..e540f11 100644 --- a/src/contextseek/config/factory.py +++ b/src/contextseek/config/factory.py @@ -13,6 +13,7 @@ from contextseek.config.settings import ( EmbeddingSettings, LLMSettings, + RetrievalSettings, SummarizerSettings, ) @@ -217,9 +218,33 @@ def build_summarizer( return None +def build_reranker(settings: RetrievalSettings) -> Any | None: + """Build the configured non-LLM reranker. + + ``None`` preserves the existing heuristic and LLM assembly paths. The + cross-encoder model itself is loaded lazily on the first retrieval. + """ + mode = settings.reranker_mode.strip().lower().replace("-", "_") + if mode in {"heuristic", "llm"}: + return None + if mode == "cross_encoder": + from contextseek.retrieval.components import CrossEncoderReranker + + return CrossEncoderReranker( + settings.cross_encoder_model, + device=settings.cross_encoder_device or None, + top_n=max(1, int(settings.cross_encoder_top_n)), + ) + raise ValueError( + f"Unknown reranker mode '{settings.reranker_mode}'. " + "Supported modes: heuristic, llm, cross_encoder." + ) + + __all__ = [ "build_embedder", "build_llm", + "build_reranker", "build_summarizer", "resolve_embedding_dims", ] diff --git a/src/contextseek/config/settings.py b/src/contextseek/config/settings.py index d40d88e..2c0b950 100644 --- a/src/contextseek/config/settings.py +++ b/src/contextseek/config/settings.py @@ -292,6 +292,9 @@ class RetrievalSettings(BaseSettings): ) reranker_mode: str = "heuristic" llm_rerank_top_n: int = 20 + cross_encoder_model: str = "cross-encoder/ms-marco-MiniLM-L-6-v2" + cross_encoder_device: str = "" + cross_encoder_top_n: int = 20 hierarchical_alpha: float = 0.5 hierarchical_max_rounds: int = 24 hierarchical_convergence_rounds: int = 3 diff --git a/src/contextseek/config/strategies.py b/src/contextseek/config/strategies.py index 8c2dc9e..7e672a1 100644 --- a/src/contextseek/config/strategies.py +++ b/src/contextseek/config/strategies.py @@ -65,7 +65,7 @@ class RetrievalStrategy: importance_floor: float = 0.1 # Geo decay: distance decay unit in km for reranker spatial penalty distance_decay_km: float = 1.0 - # Rerank mode: "heuristic" (default) or "llm" + # Rerank mode: "heuristic" (default), "llm", or "cross_encoder" reranker_mode: str = "heuristic" # Limit number of candidates scored by LLM in reranking llm_rerank_top_n: int = 20 diff --git a/src/contextseek/retrieval/__init__.py b/src/contextseek/retrieval/__init__.py index 29c1ddf..6f90eb7 100644 --- a/src/contextseek/retrieval/__init__.py +++ b/src/contextseek/retrieval/__init__.py @@ -1,6 +1,7 @@ """Retrieval pipeline exports.""" from contextseek.retrieval.components import DefaultRecallRoute +from contextseek.retrieval.components import CrossEncoderReranker from contextseek.retrieval.components import HeuristicReranker from contextseek.retrieval.components import RecallQuery from contextseek.retrieval.components import RecallRoute @@ -10,6 +11,7 @@ __all__ = [ "DefaultRecallRoute", + "CrossEncoderReranker", "HeuristicReranker", "RecallQuery", "RecallRoute", diff --git a/src/contextseek/retrieval/components.py b/src/contextseek/retrieval/components.py index 0bb1e93..99742e1 100644 --- a/src/contextseek/retrieval/components.py +++ b/src/contextseek/retrieval/components.py @@ -750,6 +750,77 @@ def rerank( return scored + remainder +class CrossEncoderReranker: + """Batch-rerank candidates with a sentence-transformers cross-encoder. + + Model loading is lazy so the default heuristic and LLM paths do not import + the optional ``sentence-transformers`` dependency. Tests and custom + integrations may inject any object exposing ``predict(pairs)``. + """ + + def __init__( + self, + model_name: str, + *, + device: str | None = None, + inner: Reranker | None = None, + top_n: int | None = None, + model: Any | None = None, + ) -> None: + self._model_name = model_name + self._device = device + self._inner = inner or HeuristicReranker() + self._top_n = top_n + self._model = model + + def _get_model(self) -> Any: + if self._model is None: + try: + from sentence_transformers import CrossEncoder + except ImportError as exc: + raise RuntimeError( + "Cross-encoder reranking requires the optional dependency; " + "install it with `pip install 'contextseek[rerank]'`." + ) from exc + kwargs = {"device": self._device} if self._device else {} + self._model = CrossEncoder(self._model_name, **kwargs) + return self._model + + def rerank( + self, + candidates: list[dict[str, object]], + *, + query: str, + strategy: RetrievalStrategy, + geo_query: Any | None = None, + ) -> list[dict[str, object]]: + pre_ranked = self._inner.rerank( + candidates, query=query, strategy=strategy, geo_query=geo_query + ) + to_score = pre_ranked[: self._top_n] if self._top_n else pre_ranked + remainder = pre_ranked[self._top_n :] if self._top_n else [] + if not to_score: + return pre_ranked + + pairs = [(query, str(item.get("content", ""))) for item in to_score] + model = self._get_model() + try: + scores = list(model.predict(pairs)) + if len(scores) != len(to_score): + raise ValueError("cross-encoder returned an unexpected score count") + for item, score in zip(to_score, scores): + item["_score"] = round(float(score), 6) + except Exception: # noqa: BLE001 + return pre_ranked + + scored = sorted( + to_score, + key=lambda item: float(item.get("_score", 0.0)), + reverse=True, + ) + return scored + remainder + + class RelationAwareReranker: """Reranker that applies relation-based boosts and penalties. diff --git a/tests/unit_tests/retrieval/test_reranker_features.py b/tests/unit_tests/retrieval/test_reranker_features.py index 44ca7fe..508dac0 100644 --- a/tests/unit_tests/retrieval/test_reranker_features.py +++ b/tests/unit_tests/retrieval/test_reranker_features.py @@ -4,7 +4,7 @@ from __future__ import annotations from contextseek.config.strategies import RetrievalStrategy -from contextseek.retrieval.components import HeuristicReranker +from contextseek.retrieval.components import CrossEncoderReranker, HeuristicReranker def _candidate(**kwargs) -> dict: @@ -13,6 +13,66 @@ def _candidate(**kwargs) -> dict: return base +class StubCrossEncoder: + def __init__(self, scores: list[float]) -> None: + self.scores = scores + self.pairs: list[tuple[str, str]] = [] + + def predict(self, pairs: list[tuple[str, str]]) -> list[float]: + self.pairs = pairs + return self.scores + + +class TestCrossEncoderReranker: + def test_batch_scores_and_reorders_candidates(self) -> None: + model = StubCrossEncoder([0.1, 0.9]) + reranker = CrossEncoderReranker("stub", model=model) + candidates = [ + _candidate(id="first", content="alpha", score=0.9, stage="skill"), + _candidate(id="second", content="beta", score=0.8, stage="skill"), + ] + + ranked = reranker.rerank( + candidates, query="question", strategy=RetrievalStrategy() + ) + + assert [item["id"] for item in ranked] == ["second", "first"] + assert model.pairs == [("question", "alpha"), ("question", "beta")] + + def test_top_n_preserves_unscored_remainder(self) -> None: + model = StubCrossEncoder([0.2, 0.8]) + reranker = CrossEncoderReranker("stub", model=model, top_n=2) + candidates = [ + _candidate(id="a", score=0.9, stage="skill"), + _candidate(id="b", score=0.8, stage="skill"), + _candidate(id="c", score=0.7, stage="skill"), + ] + + ranked = reranker.rerank( + candidates, query="q", strategy=RetrievalStrategy() + ) + + assert [item["id"] for item in ranked] == ["b", "a", "c"] + assert len(model.pairs) == 2 + + def test_prediction_failure_falls_back_to_inner_order(self) -> None: + class FailingModel: + def predict(self, pairs): + raise RuntimeError("offline") + + reranker = CrossEncoderReranker("stub", model=FailingModel()) + candidates = [ + _candidate(id="low", score=0.2, stage="skill"), + _candidate(id="high", score=0.8, stage="skill"), + ] + + ranked = reranker.rerank( + candidates, query="q", strategy=RetrievalStrategy() + ) + + assert [item["id"] for item in ranked] == ["high", "low"] + + class TestFeedbackChannel: def test_feedback_zero_does_not_bias(self) -> None: """feedback_score=0 (explicit) and absent feedback_score must produce diff --git a/tests/unit_tests/test_settings.py b/tests/unit_tests/test_settings.py index b4cdb27..b824199 100644 --- a/tests/unit_tests/test_settings.py +++ b/tests/unit_tests/test_settings.py @@ -20,6 +20,7 @@ _import_class, build_embedder, build_llm, + build_reranker, resolve_embedding_dims, ) @@ -220,6 +221,32 @@ def test_build_embedder_none(self): result = build_embedder(EmbeddingSettings()) assert result is None + def test_build_cross_encoder_reranker_lazily(self): + """Cross-encoder config builds without importing the optional package.""" + from contextseek.retrieval.components import CrossEncoderReranker + + reranker = build_reranker( + RetrievalSettings( + reranker_mode="cross_encoder", + cross_encoder_model="example/reranker", + cross_encoder_device="cpu", + cross_encoder_top_n=7, + ) + ) + + assert isinstance(reranker, CrossEncoderReranker) + assert reranker._model_name == "example/reranker" + assert reranker._device == "cpu" + assert reranker._top_n == 7 + + def test_build_reranker_rejects_unknown_mode(self): + with pytest.raises(ValueError, match="Supported modes"): + build_reranker(RetrievalSettings(reranker_mode="unknown")) + + @pytest.mark.parametrize("mode", ["heuristic", "llm"]) + def test_build_reranker_preserves_existing_modes(self, mode): + assert build_reranker(RetrievalSettings(reranker_mode=mode)) is None + def test_build_embedder_no_class_path(self): """Provider set but empty class_path returns None.""" result = build_embedder(EmbeddingSettings(provider="langchain", class_path="")) @@ -459,6 +486,20 @@ def test_from_settings_file_backend(self, tmp_path): response = ctx.retrieve("file", scope="t/p/u") assert len(response) >= 1 + def test_from_settings_selects_cross_encoder_reranker(self): + """Reranker implementation is swappable through retrieval settings.""" + from contextseek import ContextSeek + from contextseek.retrieval.components import CrossEncoderReranker + + settings = ContextSeekSettings( + storage=StorageSettings(backend="memory"), + retrieval=RetrievalSettings(reranker_mode="cross_encoder"), + ) + + ctx = ContextSeek.from_settings(settings) + + assert isinstance(ctx.reranker, CrossEncoderReranker) + def test_from_settings_with_evolution(self): """from_settings() enables evolution engine when configured.""" from contextseek import ContextSeek diff --git a/uv.lock b/uv.lock index 69da38d..212f0c2 100644 --- a/uv.lock +++ b/uv.lock @@ -497,6 +497,7 @@ all = [ { name = "langgraph" }, { name = "pydantic" }, { name = "pyseekdb" }, + { name = "sentence-transformers" }, { name = "uvicorn" }, { name = "watchdog" }, ] @@ -539,6 +540,9 @@ openai = [ powermem = [ { name = "powermem" }, ] +rerank = [ + { name = "sentence-transformers" }, +] seekdb = [ { name = "pyseekdb" }, ] @@ -590,6 +594,8 @@ requires-dist = [ { name = "pyyaml", specifier = ">=6.0" }, { name = "rich", specifier = ">=13.7.0" }, { name = "seekvfs", specifier = ">=0.1.0" }, + { name = "sentence-transformers", marker = "extra == 'all'", specifier = ">=3.0.0" }, + { name = "sentence-transformers", marker = "extra == 'rerank'", specifier = ">=3.0.0" }, { name = "sqlalchemy", marker = "extra == 'oceanbase'", specifier = ">=2.0" }, { name = "typing-extensions", specifier = ">=4.10.0" }, { name = "uvicorn", marker = "extra == 'all'", specifier = ">=0.30.0" }, @@ -597,7 +603,7 @@ requires-dist = [ { name = "watchdog", marker = "extra == 'all'", specifier = ">=4.0" }, { name = "watchdog", marker = "extra == 'daemon'", specifier = ">=4.0" }, ] -provides-extras = ["test", "http", "langchain", "openai", "dashscope", "ollama", "huggingface", "oceanbase", "daemon", "seekdb", "powermem", "all", "appworld-eval"] +provides-extras = ["test", "http", "langchain", "openai", "dashscope", "ollama", "huggingface", "oceanbase", "daemon", "seekdb", "powermem", "rerank", "all", "appworld-eval"] [package.metadata.requires-dev] dev = [ @@ -682,7 +688,7 @@ name = "cuda-bindings" version = "13.3.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "cuda-pathfinder", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "cuda-pathfinder" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/f9/52/50673d25e46d199556f827514bf646a49471d50538c5e577201245b348a9/cuda_bindings-13.3.0-cp311-cp311-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:244a167d81a9d07d3209d02c8e02e848c2b7c38f4d01e8e4d1f9620b173ae006", size = 6051409, upload-time = "2026-05-27T03:59:01.648Z" }, @@ -720,34 +726,34 @@ wheels = [ [package.optional-dependencies] cudart = [ - { name = "nvidia-cuda-runtime", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cuda-runtime", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] cufft = [ - { name = "nvidia-cufft", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cufft", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] cufile = [ - { name = "nvidia-cufile", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cufile", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] cupti = [ - { name = "nvidia-cuda-cupti", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cuda-cupti", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] curand = [ - { name = "nvidia-curand", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-curand", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] cusolver = [ - { name = "nvidia-cusolver", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cusolver", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] cusparse = [ - { name = "nvidia-cusparse", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cusparse", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] nvjitlink = [ - { name = "nvidia-nvjitlink", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nvjitlink", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] nvrtc = [ - { name = "nvidia-cuda-nvrtc", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cuda-nvrtc", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] nvtx = [ - { name = "nvidia-nvtx", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version >= '3.14' and sys_platform == 'win32') or (sys_platform != 'linux' and sys_platform != 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform == 'win32' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nvtx", marker = "sys_platform == 'linux' or sys_platform == 'win32' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, ] [[package]] @@ -1403,8 +1409,7 @@ source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ { name = "grpcio" }, { name = "protobuf", version = "6.33.6", source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } }, - { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" }, marker = "(python_full_version >= '3.14' and extra == 'group-11-contextseek-dev') or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "setuptools", version = "82.0.1", source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" }, marker = "(python_full_version < '3.14' and extra == 'group-11-contextseek-dev') or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "setuptools" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/8b/d1/cbefe328653f746fd319c4377836a25ba64226e41c6a1d7d5cdbc87a459f/grpcio_tools-1.78.0.tar.gz", hash = "sha256:4b0dd86560274316e155d925158276f8564508193088bc43e20d3f5dff956b2b", size = 5393026, upload-time = "2026-02-06T09:59:59.53Z" } wheels = [ @@ -1632,7 +1637,7 @@ name = "jinja2" version = "3.1.6" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "markupsafe", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "markupsafe" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/df/bf/f7da0350254c0ed7c72f3e33cef02e048281fec7ecec5f032d4aac52226b/jinja2-3.1.6.tar.gz", hash = "sha256:0137fb05990d35f1275a587e9aee6d56da821fc83491a0fb838183be43f66d6d", size = 245115, upload-time = "2025-03-05T20:05:02.478Z" } wheels = [ @@ -2451,7 +2456,7 @@ name = "nvidia-cublas" version = "13.1.1.3" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "nvidia-cuda-nvrtc", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cuda-nvrtc" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/a7/a1/0bd24ee8c8d03adac032fd2909426a00c88f8c57961b1277ded97f91119f/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:b7a210458267ac818974c53038fbec2e969d5c99f305ab15c72522fa9f001dd5", size = 542848918, upload-time = "2026-04-08T18:46:22.985Z" }, @@ -2494,7 +2499,7 @@ name = "nvidia-cudnn-cu13" version = "9.20.0.48" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "nvidia-cublas", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cublas" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/56/c5/83384d846b2fd17c44bd499b36c75a45ed4f095fbbb2252294e89cea5c5c/nvidia_cudnn_cu13-9.20.0.48-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:e31454ae00094b0c55319d9d15b6fa2fc50a9e1c0f5c8c80fb75258234e731e1", size = 444574296, upload-time = "2026-03-09T19:28:27.751Z" }, @@ -2507,7 +2512,7 @@ name = "nvidia-cufft" version = "12.0.0.61" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "nvidia-nvjitlink", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nvjitlink" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/8b/ae/f417a75c0259e85c1d2f83ca4e960289a5f814ed0cea74d18c353d3e989d/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:2708c852ef8cd89d1d2068bdbece0aa188813a0c934db3779b9b1faa8442e5f5", size = 214053554, upload-time = "2025-09-04T08:31:38.196Z" }, @@ -2539,9 +2544,9 @@ name = "nvidia-cusolver" version = "12.0.4.66" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "nvidia-cublas", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-cusparse", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-nvjitlink", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cublas" }, + { name = "nvidia-cusparse" }, + { name = "nvidia-nvjitlink" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/c8/c3/b30c9e935fc01e3da443ec0116ed1b2a009bb867f5324d3f2d7e533e776b/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:02c2457eaa9e39de20f880f4bd8820e6a1cfb9f9a34f820eb12a155aa5bc92d2", size = 223467760, upload-time = "2025-09-04T08:33:04.222Z" }, @@ -2554,7 +2559,7 @@ name = "nvidia-cusparse" version = "12.6.3.3" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "nvidia-nvjitlink", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nvjitlink" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/f8/94/5c26f33738ae35276672f12615a64bd008ed5be6d1ebcb23579285d960a9/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:80bcc4662f23f1054ee334a15c72b8940402975e0eab63178fc7e670aa59472c", size = 162155568, upload-time = "2025-09-04T08:33:42.864Z" }, @@ -3910,10 +3915,10 @@ name = "scikit-learn" version = "1.8.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "joblib", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "numpy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "scipy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "threadpoolctl", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "joblib" }, + { name = "numpy" }, + { name = "scipy" }, + { name = "threadpoolctl" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/0e/d4/40988bf3b8e34feec1d0e6a051446b1f66225f8529b9309becaeef62b6c4/scikit_learn-1.8.0.tar.gz", hash = "sha256:9bccbb3b40e3de10351f8f5068e105d0f4083b1a65fa07b6634fbc401a6287fd", size = 7335585, upload-time = "2025-12-10T07:08:53.618Z" } wheels = [ @@ -3960,7 +3965,7 @@ name = "scipy" version = "1.17.1" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "numpy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "numpy" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/7a/97/5a3609c4f8d58b039179648e62dd220f89864f56f7357f5d4f45c29eb2cc/scipy-1.17.1.tar.gz", hash = "sha256:95d8e012d8cb8816c226aef832200b1d45109ed4464303e997c5b13122b297c0", size = 30573822, upload-time = "2026-02-23T00:26:24.851Z" } wheels = [ @@ -4044,14 +4049,14 @@ name = "sentence-transformers" version = "5.5.1" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "huggingface-hub", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "numpy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "scikit-learn", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "scipy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "torch", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "tqdm", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "transformers", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "typing-extensions", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "huggingface-hub" }, + { name = "numpy" }, + { name = "scikit-learn" }, + { name = "scipy" }, + { name = "torch" }, + { name = "tqdm" }, + { name = "transformers" }, + { name = "typing-extensions" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/cf/d4/7ef93157485e978c016f49da05363c1e4e7237beb5343b64b5631101f0f1/sentence_transformers-5.5.1.tar.gz", hash = "sha256:02b7740dfc60bdbbcb6061625f5d97a5c1a4e2d3baac5f9391b912bb5eae2290", size = 445161, upload-time = "2026-05-20T07:37:44.465Z" } wheels = [ @@ -4062,27 +4067,11 @@ wheels = [ name = "setuptools" version = "81.0.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } -resolution-markers = [ - "python_full_version >= '3.14'", -] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/0d/1c/73e719955c59b8e424d015ab450f51c0af856ae46ea2da83eba51cc88de1/setuptools-81.0.0.tar.gz", hash = "sha256:487b53915f52501f0a79ccfd0c02c165ffe06631443a886740b91af4b7a5845a", size = 1198299, upload-time = "2026-02-06T21:10:39.601Z" } wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/e1/e3/c164c88b2e5ce7b24d667b9bd83589cf4f3520d97cad01534cd3c4f55fdb/setuptools-81.0.0-py3-none-any.whl", hash = "sha256:fdd925d5c5d9f62e4b74b30d6dd7828ce236fd6ed998a08d81de62ce5a6310d6", size = 1062021, upload-time = "2026-02-06T21:10:37.175Z" }, ] -[[package]] -name = "setuptools" -version = "82.0.1" -source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } -resolution-markers = [ - "python_full_version == '3.13.*'", - "python_full_version < '3.13'", -] -sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/4f/db/cfac1baf10650ab4d1c111714410d2fbb77ac5a616db26775db562c8fab2/setuptools-82.0.1.tar.gz", hash = "sha256:7d872682c5d01cfde07da7bccc7b65469d3dca203318515ada1de5eda35efbf9", size = 1152316, upload-time = "2026-03-09T12:47:17.221Z" } -wheels = [ - { url = "https://pypi.tuna.tsinghua.edu.cn/packages/9d/76/f789f7a86709c6b087c5a2f52f911838cad707cc613162401badc665acfe/setuptools-82.0.1-py3-none-any.whl", hash = "sha256:a59e362652f08dcd477c78bb6e7bd9d80a7995bc73ce773050228a348ce2e5bb", size = 1006223, upload-time = "2026-03-09T12:47:15.026Z" }, -] - [[package]] name = "shapely" version = "2.1.2" @@ -4302,7 +4291,7 @@ name = "sympy" version = "1.14.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "mpmath", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "mpmath" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/83/d3/803453b36afefb7c2bb238361cd4ae6125a569b4db67cd9e79846ba2d68c/sympy-1.14.0.tar.gz", hash = "sha256:d3d3fe8df1e5a0b42f0e7bdf50541697dbe7d23746e894990c030e2b05e72517", size = 7793921, upload-time = "2025-04-27T18:05:01.611Z" } wheels = [ @@ -4448,21 +4437,21 @@ name = "torch" version = "2.12.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "cuda-bindings", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "cuda-toolkit", extra = ["cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "filelock", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "fsspec", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "jinja2", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "networkx", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-cublas", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-cudnn-cu13", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-cusparselt-cu13", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-nccl-cu13", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "nvidia-nvshmem-cu13", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" }, marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "sympy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "triton", marker = "(python_full_version >= '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev') or (sys_platform != 'linux' and extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "typing-extensions", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "cuda-bindings", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "cuda-toolkit", extra = ["cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "filelock" }, + { name = "fsspec" }, + { name = "jinja2" }, + { name = "networkx" }, + { name = "nvidia-cublas", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cudnn-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-cusparselt-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nccl-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "nvidia-nvshmem-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "setuptools" }, + { name = "sympy" }, + { name = "triton", marker = "sys_platform == 'linux' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "typing-extensions" }, ] wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/18/62/131124fb95df03811b8260d1d43dcc5ee85ea1a344b964613d7efe77fb08/torch-2.12.0-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:10802fd383bbfed646212e765a72c37d2185205d4f26eb197a254e8ac7ddcb25", size = 87990344, upload-time = "2026-05-13T14:55:42.154Z" }, @@ -4508,15 +4497,15 @@ name = "transformers" version = "5.9.0" source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } dependencies = [ - { name = "huggingface-hub", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "numpy", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "packaging", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "pyyaml", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "regex", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "safetensors", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "tokenizers", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "tqdm", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, - { name = "typer", marker = "python_full_version >= '3.14' or (extra == 'extra-11-contextseek-powermem' and extra == 'group-11-contextseek-dev')" }, + { name = "huggingface-hub" }, + { name = "numpy" }, + { name = "packaging" }, + { name = "pyyaml" }, + { name = "regex" }, + { name = "safetensors" }, + { name = "tokenizers" }, + { name = "tqdm" }, + { name = "typer" }, ] sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/51/58/7f843608f2e8421f86bb97060b54649be6239ec612b82bf9d41e65c26c00/transformers-5.9.0.tar.gz", hash = "sha256:25997cb8fa6053533171634b6162d7df54346530ec2aa9b42bb834e63668c842", size = 8642240, upload-time = "2026-05-20T14:50:49.278Z" } wheels = [ From cca5d08f6e1386030ea919e75b87fa7ed4282bac Mon Sep 17 00:00:00 2001 From: vcvyg <180424490@qq.com> Date: Fri, 31 Jul 2026 18:22:05 +0800 Subject: [PATCH 2/2] fix: format reranker test to pass ruff format check --- tests/unit_tests/retrieval/test_reranker_features.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/tests/unit_tests/retrieval/test_reranker_features.py b/tests/unit_tests/retrieval/test_reranker_features.py index 508dac0..2116ec0 100644 --- a/tests/unit_tests/retrieval/test_reranker_features.py +++ b/tests/unit_tests/retrieval/test_reranker_features.py @@ -48,9 +48,7 @@ def test_top_n_preserves_unscored_remainder(self) -> None: _candidate(id="c", score=0.7, stage="skill"), ] - ranked = reranker.rerank( - candidates, query="q", strategy=RetrievalStrategy() - ) + ranked = reranker.rerank(candidates, query="q", strategy=RetrievalStrategy()) assert [item["id"] for item in ranked] == ["b", "a", "c"] assert len(model.pairs) == 2 @@ -66,9 +64,7 @@ def predict(self, pairs): _candidate(id="high", score=0.8, stage="skill"), ] - ranked = reranker.rerank( - candidates, query="q", strategy=RetrievalStrategy() - ) + ranked = reranker.rerank(candidates, query="q", strategy=RetrievalStrategy()) assert [item["id"] for item in ranked] == ["high", "low"]