feat(python): added support for WatsonxReranker component (#3642)

## Summary

Adds `WatsonxReranker` to the Python bindings, integrating the [IBM
watsonx.ai text rerank
API](https://cloud.ibm.com/docs/apis/watsonx-ai#text-rerank) via the
`ibm_watsonx_ai` SDK (`pip install ibm-watsonx-ai`).

## Parameters

| Parameter | Default | Description |
|---|---|---|
| `model_name` | `"cross-encoder/ms-marco-minilm-l-12-v2"` | Rerank
model ID |
| `column` | `"text"` | Table column used as document input |
| `top_n` | `None` | Return only the top-n results |
| `return_score` | `"relevance"` | `"relevance"` or `"all"` |
| `api_key` | `None` | Falls back to `WATSONX_API_KEY` env var |
| `project_id` | `None` | Falls back to `WATSONX_PROJECT_ID` env var —
mutually exclusive with `space_id` |
| `space_id` | `None` | Falls back to `WATSONX_SPACE_ID` env var —
mutually exclusive with `project_id` |
| `url` | `None` | Defaults to `https://us-south.ml.cloud.ibm.com` |
| `truncate_input_tokens` | `None` | Token truncation limit |

## Usage

```python
from lancedb.rerankers import WatsonxReranker

# credentials from environment variables
reranker = WatsonxReranker()

# or passed explicitly
reranker = WatsonxReranker(
    api_key="<key>",
    project_id="<project-id>",   # or space_id="<space-id>"
    top_n=5,
)
```

## Testing

Integration test added in `test_rerankers.py`, skipped unless
`WATSONX_API_KEY` and one of `WATSONX_PROJECT_ID` / `WATSONX_SPACE_ID`
are set.
This commit is contained in:
Mateusz Szewczyk
2026-07-14 00:58:32 +02:00
committed by GitHub
parent cde48fad95
commit 5b982f2f05
3 changed files with 199 additions and 0 deletions
@@ -12,6 +12,7 @@ from .rrf import RRFReranker
from .mrr import MRRReranker
from .answerdotai import AnswerdotaiRerankers
from .voyageai import VoyageAIReranker
from .watsonx import WatsonxReranker
__all__ = [
"Reranker",
@@ -25,4 +26,5 @@ __all__ = [
"AnswerdotaiRerankers",
"VoyageAIReranker",
"MRRReranker",
"WatsonxReranker",
]
+180
View File
@@ -0,0 +1,180 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import os
from functools import cached_property
from typing import Dict, Optional
import pyarrow as pa
from ..util import attempt_import_or_raise
from .base import Reranker
DEFAULT_WATSONX_URL = "https://us-south.ml.cloud.ibm.com"
class WatsonxReranker(Reranker):
"""
Reranks the results using the IBM watsonx.ai Rerank API.
Uses the ``ibm_watsonx_ai`` SDK (``Rerank.generate``) under the hood.
API Docs:
https://cloud.ibm.com/docs/apis/watsonx-ai#text-rerank
Supported rerank models:
https://dataplatform.cloud.ibm.com/docs/content/wsj/analyze-data/fm-models-embed.html?context=wx#rerank
Parameters
----------
model_name : str, default "cross-encoder/ms-marco-minilm-l-12-v2"
The ID of the rerank model to use.
column : str, default "text"
The name of the column to use as input to the reranker.
top_n : int, optional
Return only the top-n results. If ``None``, all results are returned.
return_score : str, default "relevance"
Options are ``"relevance"`` or ``"all"``.
api_key : str, optional
IBM Cloud API key. Falls back to the ``WATSONX_API_KEY`` environment
variable when not provided.
project_id : str, optional
watsonx.ai project ID. Falls back to the ``WATSONX_PROJECT_ID``
environment variable when not provided. Mutually exclusive with
``space_id`` — exactly one must be supplied.
space_id : str, optional
watsonx.ai deployment space ID. Falls back to the ``WATSONX_SPACE_ID``
environment variable when not provided. Mutually exclusive with
``project_id`` — exactly one must be supplied.
url : str, optional
watsonx.ai service URL. Defaults to
``"https://us-south.ml.cloud.ibm.com"``.
truncate_input_tokens : int, optional
Truncate each input to this many tokens before scoring. Passed
directly to the ``parameters`` dict of ``Rerank.generate``.
"""
def __init__(
self,
model_name: str = "cross-encoder/ms-marco-minilm-l-12-v2",
column: str = "text",
top_n: Optional[int] = None,
return_score: str = "relevance",
api_key: Optional[str] = None,
project_id: Optional[str] = None,
space_id: Optional[str] = None,
url: Optional[str] = None,
truncate_input_tokens: Optional[int] = None,
):
super().__init__(return_score)
self.model_name = model_name
self.column = column
self.top_n = top_n
self.api_key = api_key
self.project_id = project_id
self.space_id = space_id
self.url = url
self.truncate_input_tokens = truncate_input_tokens
def __str__(self) -> str:
return f"WatsonxReranker(model_name={self.model_name})"
@cached_property
def _client(self):
ibm_watsonx_ai = attempt_import_or_raise("ibm_watsonx_ai")
ibm_watsonx_ai_foundation_models = attempt_import_or_raise(
"ibm_watsonx_ai.foundation_models"
)
# --- credentials ---
api_key = self.api_key or os.environ.get("WATSONX_API_KEY")
if not api_key:
raise ValueError(
"WATSONX_API_KEY not set. Either set it in your environment or "
"pass it as `api_key` argument to WatsonxReranker."
)
credentials = ibm_watsonx_ai.Credentials(
api_key=api_key,
url=self.url or DEFAULT_WATSONX_URL,
)
# --- project_id / space_id (exactly one required) ---
project_id = self.project_id or os.environ.get("WATSONX_PROJECT_ID")
space_id = self.space_id or os.environ.get("WATSONX_SPACE_ID")
if project_id and space_id:
raise ValueError("Provide either `project_id` or `space_id`, not both.")
if not project_id and not space_id:
raise ValueError(
"Either WATSONX_PROJECT_ID or WATSONX_SPACE_ID must be set. "
"Pass one as an argument to WatsonxReranker or set the corresponding "
"environment variable."
)
kwargs: Dict = dict(model_id=self.model_name, credentials=credentials)
if project_id:
kwargs["project_id"] = project_id
else:
kwargs["space_id"] = space_id
return ibm_watsonx_ai_foundation_models.Rerank(**kwargs)
def _build_params(self) -> Dict:
"""Build the ``parameters`` dict forwarded to ``Rerank.generate``."""
return_options: Dict = {"inputs": True}
if self.top_n is not None:
return_options["top_n"] = self.top_n
params: Dict = {"return_options": return_options}
if self.truncate_input_tokens is not None:
params["truncate_input_tokens"] = self.truncate_input_tokens
return params
def _rerank(self, result_set: pa.Table, query: str) -> pa.Table:
result_set = self._handle_empty_results(result_set)
if len(result_set) == 0:
return result_set
docs = result_set[self.column].to_pylist()
response = self._client.generate(
query=query,
inputs=docs,
params=self._build_params(),
)
results = response["results"]
indices, scores = zip(
*[(result["index"], result["score"]) for result in results]
)
result_set = result_set.take(list(indices))
result_set = result_set.append_column(
"_relevance_score", pa.array(scores, type=pa.float32())
)
return result_set
def rerank_hybrid(
self,
query: str,
vector_results: pa.Table,
fts_results: pa.Table,
) -> pa.Table:
if self.score == "all":
combined_results = self._merge_and_keep_scores(vector_results, fts_results)
else:
combined_results = self.merge_results(vector_results, fts_results)
combined_results = self._rerank(combined_results, query)
if self.score == "relevance":
combined_results = self._keep_relevance_score(combined_results)
return combined_results
def rerank_vector(self, query: str, vector_results: pa.Table) -> pa.Table:
vector_results = self._rerank(vector_results, query)
if self.score == "relevance":
vector_results = vector_results.drop_columns(["_distance"])
return vector_results
def rerank_fts(self, query: str, fts_results: pa.Table) -> pa.Table:
fts_results = self._rerank(fts_results, query)
if self.score == "relevance":
fts_results = fts_results.drop_columns(["_score"])
return fts_results
+17
View File
@@ -23,6 +23,7 @@ from lancedb.rerankers import (
AnswerdotaiRerankers,
VoyageAIReranker,
MRRReranker,
WatsonxReranker,
)
from lancedb.table import LanceTable
@@ -727,3 +728,19 @@ def test_linear_combination_missing_fts_is_penalised():
f"Document with FTS score (rowid 0, {scores[0]:.4f}) should beat "
f"document with no FTS match (rowid 1, {scores[1]:.4f})"
)
@pytest.mark.skipif(
os.environ.get("WATSONX_API_KEY") is None
or (
os.environ.get("WATSONX_PROJECT_ID") is None
and os.environ.get("WATSONX_SPACE_ID") is None
),
reason="WATSONX_API_KEY and one of WATSONX_PROJECT_ID / "
"WATSONX_SPACE_ID must be set",
)
def test_watsonx_reranker(tmp_path):
pytest.importorskip("ibm_watsonx_ai")
table, schema = get_test_table(tmp_path)
reranker = WatsonxReranker()
_run_test_reranker(reranker, table, "single player experience", None, schema)