From 5062a77f17d0f98e9c11a17ee723afdfd7f08541 Mon Sep 17 00:00:00 2001 From: chaoyang Date: Sun, 20 Sep 2026 16:58:26 +0800 Subject: [PATCH] [python] Add search result scores and explicit relevance ordering --- docs/docs/pypaimon/multimodal-search.md | 51 +++++ paimon-python/pypaimon/multimodal/query.py | 99 ++++++++- .../tests/ray_search_result_metadata_test.py | 57 +++++ .../tests/search_result_metadata_test.py | 206 ++++++++++++++++++ 4 files changed, 407 insertions(+), 6 deletions(-) create mode 100644 paimon-python/pypaimon/tests/ray_search_result_metadata_test.py create mode 100644 paimon-python/pypaimon/tests/search_result_metadata_test.py diff --git a/docs/docs/pypaimon/multimodal-search.md b/docs/docs/pypaimon/multimodal-search.md index 8c98e2e3313e..05d38be456d9 100644 --- a/docs/docs/pypaimon/multimodal-search.md +++ b/docs/docs/pypaimon/multimodal-search.md @@ -109,6 +109,57 @@ matches = ( ) ``` +## Scores and Result Ordering + +On data-evolution tables, use `with_score()` to append a `float64` relevance +column and `order_by_score()` to sort by descending score, with ascending +`_ROW_ID` for ties. Both methods are optional: `with_score()` alone preserves +the existing result order, and `order_by_score()` does not require projecting +scores. Without either method, result behavior is unchanged. + +```python +neighbors = ( + docs.search([0.1, 0.2, 0.3], column="embedding") + .select(["id", "content"]) + .with_score("relevance") + .order_by_score() + .limit(10) + .to_arrow() +) +``` + +The default score column is `_score`. A custom name must not conflict with a +table column or a system field. Scores use the search engine's existing +higher-is-better convention: L2 uses `1 / (1 + squared_distance)`, cosine uses +cosine similarity, and inner product uses the dot product. Full-text results +expose BM25 scores; hybrid results expose the selected ranker's fusion scores. +Scores from different metrics or rankers are not directly comparable. + +These methods also work with full-text, hybrid, and batch vector queries, and +with local or Ray vector execution. Batch output retains input-query order and +each row receives its score for that query. `where()` still filters selected +rows during lookup, so it can return fewer than the requested number of hits. + +When only row IDs and scores are needed, explicitly project `_ROW_ID`: + +```python +hits = ( + docs.search([0.1, 0.2, 0.3], column="embedding") + .select(["_ROW_ID"]) + .with_score() + .order_by_score() + .limit(10) + .to_arrow() +) +``` + +Use `select([]).with_score()` for scores alone. When either score method is +enabled and the explicit projection contains only `_ROW_ID` or is empty, the +query skips final row lookup if there is no `where()` and query authorization +is disabled. Raw search, prefiltering, and vector refinement can still read +data. Historical snapshot and deletion semantics remain the same. Plain +`select([])` without either method retains its existing behavior. + ## Distributed Vector Search Use `execution="ray"` to execute vector queries across Ray workers and return diff --git a/paimon-python/pypaimon/multimodal/query.py b/paimon-python/pypaimon/multimodal/query.py index af2a983f9a06..a7c8a569a269 100644 --- a/paimon-python/pypaimon/multimodal/query.py +++ b/paimon-python/pypaimon/multimodal/query.py @@ -483,10 +483,96 @@ def __init__( result_factory: Optional[Callable] = None, pre_filter=None): self._pre_filter = None + self._score_column = None + self._sort_by_score = False super().__init__(table, result_factory=result_factory) if pre_filter is not None: self.pre_filter(pre_filter) + def with_score(self, column_name="_score"): + """Append relevance scores (higher is better) to data-evolution results.""" + self._require_row_id_scores() + if not isinstance(column_name, str) or not column_name: + raise ValueError("Score column name must be a nonempty string.") + if (column_name in SpecialFields.SYSTEM_FIELD_NAMES + or column_name in {field.name for field in self._table.fields}): + raise ValueError("Score column name conflicts with a table column: %s" % column_name) + self._score_column = column_name + return self + + def order_by_score(self): + """Return highest scores first, breaking ties by ascending row ID.""" + self._require_row_id_scores() + self._sort_by_score = True + return self + + def _require_row_id_scores(self): + if not self._table.options.data_evolution_enabled(): + raise NotImplementedError("Search score output and ordering require a data-evolution table.") + + def _metadata_only_result(self): + projection = self._effective_projection() + return (bool(self._score_column or self._sort_by_score) + and self._table.options.data_evolution_enabled() + and not self._table.options.query_auth_enabled + and self._predicate is None + and projection is not None + and all(name == SpecialFields.ROW_ID.name for name in projection)) + + def _read_global_index_result(self, result): + metadata_only = self._metadata_only_result() + if not (metadata_only or self._score_column or self._sort_by_score): + return super()._read_global_index_result(result) + + row_id_name = SpecialFields.ROW_ID.name + projection = self._effective_projection() + added_row_id = projection is None or row_id_name not in projection + if metadata_only: + row_ids = pa.array(list(result.results()), type=pa.int64()) + fields = [pa.field(row_id_name, pa.int64(), nullable=False)] * (len(projection) or 1) + table = pa.Table.from_arrays( + [row_ids] * len(fields), schema=pa.schema(fields)) + else: + lookup = copy(self) + lookup._projection = (list(projection) if projection is not None + else [field.name for field in self._table.fields]) + if added_row_id: + lookup._projection.append(row_id_name) + table = lookup._read_search_rows(result) + return self._finish_search_result(table, result, added_row_id) + + def _read_search_rows(self, result): + projection = self._effective_projection() + if len(projection) == len(set(projection)): + return ScanQuery._read_global_index_result(self, result) + # Row tracking requires unique names while reading. Restore repeated + # output columns after reading their values once. + fields = self._configured_read_builder().read_type() + lookup = copy(self) + lookup._projection = list(dict.fromkeys(projection)) + table = ScanQuery._read_global_index_result(lookup, result) + return table.select([table.column_names.index(field.name) for field in fields]) + + def _finish_search_result(self, table, result, added_row_id): + if self._score_column or self._sort_by_score: + row_ids = table.column(table.column_names.index(SpecialFields.ROW_ID.name)).to_pylist() + scores = [] + if row_ids: + getter = result.score_getter() + scores = [getter(row_id) for row_id in row_ids] + if any(score is None for score in scores): + raise ValueError("Missing score for a selected search row.") + if self._score_column: + if self._score_column in table.column_names: + raise ValueError("Score column name conflicts with a projected column: %s" % self._score_column) + table = table.append_column(self._score_column, pa.array(scores, type=pa.float64())) + if self._sort_by_score: + order = sorted(range(len(row_ids)), key=lambda i: (-scores[i], row_ids[i])) + table = table.take(pa.array(order, type=pa.int64())) + if added_row_id: + table = table.select([i for i, name in enumerate(table.column_names) if name != SpecialFields.ROW_ID.name]) + return table + def pre_filter(self, predicate): predicate = self._coerce_predicate(predicate, "pre_filter()") if predicate is not None: @@ -697,7 +783,8 @@ def _read_batch_results(self, results): from pypaimon.globalindex.global_index_result import GlobalIndexResult from pypaimon.utils.roaring_bitmap import RoaringBitmap64 - if len(results) <= 1 or not self._configured_read_builder().read_type(): + if (len(results) <= 1 or self._metadata_only_result() + or not self._configured_read_builder().read_type()): return [self._read_global_index_result(result) for result in results] row_ids = RoaringBitmap64() @@ -709,23 +796,23 @@ def _read_batch_results(self, results): # to the union; where() still filters the selected rows during lookup. lookup._limit = None projection = self._effective_projection() - lookup._projection = list(projection) if projection else [f.name for f in self._table.fields] + lookup._projection = (list(projection) if projection is not None and ( + projection or self._score_column or self._sort_by_score) else [f.name for f in self._table.fields]) added_row_id = SpecialFields.ROW_ID.name not in lookup._projection if added_row_id: lookup._projection.append(SpecialFields.ROW_ID.name) fields = lookup._configured_read_builder().read_type() row_id_column = next(i for i, field in enumerate(fields) if field.id == SpecialFields.ROW_ID.id) - table = lookup._read_global_index_result(GlobalIndexResult.create(row_ids)) + table = lookup._read_search_rows(GlobalIndexResult.create(row_ids)) positions = {row_id: i for i, row_id in enumerate(table.column(row_id_column).to_pylist())} - if added_row_id: - table = table.select(list(range(table.num_columns - 1))) output = [] for result in results: # Keep the physical read order, rather than imposing score or row-id order. selected = sorted(positions[row_id] for row_id in result.results() if row_id in positions) if self._limit is not None: selected = selected[:self._limit] - output.append(table.take(pa.array(selected, type=pa.int64()))) + selected_table = table.take(pa.array(selected, type=pa.int64())) + output.append(self._finish_search_result(selected_table, result, added_row_id)) return output def to_pandas(self, *, execution="local", concurrency=None, ray_remote_args=None): diff --git a/paimon-python/pypaimon/tests/ray_search_result_metadata_test.py b/paimon-python/pypaimon/tests/ray_search_result_metadata_test.py new file mode 100644 index 000000000000..570bfc3e86bc --- /dev/null +++ b/paimon-python/pypaimon/tests/ray_search_result_metadata_test.py @@ -0,0 +1,57 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +from unittest.mock import patch + +import pytest + +pytest.importorskip("ray") + +from pypaimon.multimodal.query import ScanQuery +from pypaimon.tests import ray_vector_search_test as ray_fixtures +from pypaimon.tests import search_result_metadata_test as fixtures + +ray_cluster = ray_fixtures.ray_cluster +docs = fixtures.docs + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("indexed", [False, True]) +@pytest.mark.parametrize("metadata_only", [False, True]) +def test_ray_metadata_and_ordering_match_local(docs, ray_cluster, batch, indexed, metadata_only): + if indexed: + pytest.importorskip("paimon_vindex") + docs.raw_table.copy({"deletion-vectors.enabled": "false"}).create_global_index( + "embedding", "ivf-flat", options={"ivf-flat.nlist": "1", "ivf-flat.distance.metric": "l2"}) + search = fixtures.query(docs, batch, options={"ivf-flat.refine-factor": "2"}) + search.select(["_ROW_ID"] if metadata_only else ["id"]).with_score().order_by_score() + expected = search.to_arrow() + original = ScanQuery._read_global_index_result + calls = [] + + def lookup(query, result): + calls.append(True) + return original(query, result) + + with patch.object(ScanQuery, "_read_global_index_result", lookup): + actual = search.to_arrow(execution="ray", concurrency=2) + assert len(calls) == (0 if metadata_only else 1) + expected_rows = [table.to_pylist() for table in expected] if batch else expected.to_pylist() + actual_rows = [table.to_pylist() for table in actual] if batch else actual.to_pylist() + assert actual_rows == expected_rows + as_list = search.to_list(execution="ray", concurrency=2) + assert as_list == expected_rows diff --git a/paimon-python/pypaimon/tests/search_result_metadata_test.py b/paimon-python/pypaimon/tests/search_result_metadata_test.py new file mode 100644 index 000000000000..36fb34f413c3 --- /dev/null +++ b/paimon-python/pypaimon/tests/search_result_metadata_test.py @@ -0,0 +1,206 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +from unittest.mock import patch + +import pyarrow as pa +import pytest + +import pypaimon.multimodal as pm +from pypaimon.multimodal.query import ScanQuery + + +@pytest.fixture +def docs(tmp_path): + schema = pa.schema([("id", pa.int64()), ("text", pa.string()), + ("embedding", pa.list_(pa.float32(), 2))]) + table = pm.connect(options={"warehouse": str(tmp_path)}).create_table( + "docs", schema=schema, options={"file.format": "parquet", "vector.file.format": "parquet", + "vector-index.search-mode": "full", "read.batch-size": "1"}) + table.add(pa.table({"id": [0, 1, 2, 3], "text": ["paimon long document", "paimon", "other", "paimon"], + "embedding": [[2., 1.], [0., 1.], [1., 1.], [0., 1.]]}, schema=schema)) + return table + + +def query(docs, batch=False, **kwargs): + return (docs.search_vectors([[0., 1.], [2., 1.]], **kwargs) if batch else + docs.search([0., 1.], **kwargs)).limit(4) + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("indexed", [False, True]) +def test_scores_align_with_rows_and_ordering_is_opt_in(docs, batch, indexed): + if indexed: + pytest.importorskip("paimon_vindex") + docs.raw_table.copy({"deletion-vectors.enabled": "false"}).create_global_index( + "embedding", "ivf-flat", options={"ivf-flat.nlist": "1", "ivf-flat.distance.metric": "l2"}) + plain = query(docs, batch).select(["id"]).to_arrow() + scored = query(docs, batch).select(["id"]).with_score().to_arrow() + ordered = query(docs, batch).select(["id"]).with_score("relevance").order_by_score().to_arrow() + plain, scored, ordered = ([plain], [scored], [ordered]) if not batch else (plain, scored, ordered) + expected = [[.2, 1., .5, 1.], [1., .2, .5, .2]] + for i, (before, after, ranked) in enumerate(zip(plain, scored, ordered)): + assert before.equals(after.select(["id"])) + assert after.column("_score").type == pa.float64() + expected_scores = [expected[i][row] for row in after["id"].to_pylist()] + assert after.column("_score").to_pylist() == pytest.approx(expected_scores) + ids = ranked["id"].to_pylist() + assert ids == sorted(range(4), key=lambda row: (-expected[i][row], row)) + assert ranked.column_names == ["id", "relevance"] + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("projection", [[], ["_ROW_ID"]]) +@pytest.mark.parametrize("indexed", [False, True]) +def test_metadata_only_skips_final_lookup(docs, batch, projection, indexed): + if indexed: + pytest.importorskip("paimon_vindex") + docs.raw_table.copy({"deletion-vectors.enabled": "false"}).create_global_index( + "embedding", "ivf-flat", options={"ivf-flat.nlist": "1", "ivf-flat.distance.metric": "l2"}) + with patch.object(ScanQuery, "_read_global_index_result", side_effect=AssertionError("final lookup")): + actual = query(docs, batch).select(projection).with_score().order_by_score().to_arrow() + for result in actual if batch else [actual]: + assert result.num_rows == 4 + assert result.column_names == projection + ["_score"] + assert result["_score"].to_pylist() == sorted(result["_score"].to_pylist(), reverse=True) + + +def test_ordering_does_not_require_score_projection(docs): + actual = query(docs).select(["id"]).order_by_score().to_list() + assert actual == [{"id": 1}, {"id": 3}, {"id": 2}, {"id": 0}] + + +@pytest.mark.parametrize("batch", [False, True]) +def test_score_only_with_post_filter(docs, batch): + actual = query(docs, batch).select([]).where("id >= 2").with_score().order_by_score().to_arrow() + for i, table in enumerate(actual if batch else [actual]): + expected = [1., .5] if i == 0 else [.5, .2] + assert table.column_names == ["_score"] + assert table["_score"].to_pylist() == pytest.approx(expected) + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("column", ["id", "_ROW_ID"]) +def test_duplicate_projections_keep_scores(docs, batch, column): + actual = query(docs, batch).select([column, column]).with_score().order_by_score().to_arrow() + for table in actual if batch else [actual]: + assert table.column_names == [column, column, "_score"] + assert table.column(0).equals(table.column(1)) + if column == "_ROW_ID": + with_lookup = query(docs, batch).select([column, column]).where("id >= 0").with_score().to_arrow() + assert (with_lookup[0] if batch else with_lookup).schema == (actual[0] if batch else actual).schema + + +@pytest.mark.parametrize("batch", [False, True]) +def test_post_filter_and_deleted_rows_keep_score_alignment(docs, batch): + docs.delete("id = 1") + actual = query(docs, batch).select(["id"]).where("id >= 2").with_score().order_by_score().to_arrow() + tables = actual if batch else [actual] + assert tables[0].to_pylist() == [{"id": 3, "_score": 1.}, {"id": 2, "_score": .5}] + with patch.object(ScanQuery, "_read_global_index_result", wraps=None) as lookup: + lookup.side_effect = RuntimeError("post filter needs lookup") + with pytest.raises(RuntimeError, match="post filter needs lookup"): + query(docs, batch).select(["_ROW_ID"]).where("id >= 2").with_score().to_arrow() + + +def test_query_authorization_keeps_final_lookup(docs): + from pypaimon.catalog.filesystem_catalog import FileSystemCatalog + from pypaimon.catalog.table_query_auth import TableQueryAuthResult + from pypaimon.globalindex.vector_search_result import DictBasedScoredIndexResult + from pypaimon.tests.table_query_auth_test import _simple_filter_json + + docs.raw_table = docs.raw_table.copy({"query-auth.enabled": "true"}) + result = DictBasedScoredIndexResult({0: .2, 1: 1., 2: .5, 3: 1.}) + auth = TableQueryAuthResult([_simple_filter_json("text", "paimon")], None) + search = query(docs).select(["_ROW_ID"]).with_score().order_by_score()._for_execution() + with patch.object(FileSystemCatalog, "auth_table_query", return_value=auth): + assert search._read_global_index_result(result).to_pylist() == [ + {"_ROW_ID": 1, "_score": 1.}, {"_ROW_ID": 3, "_score": 1.}] + + +@pytest.mark.parametrize("batch", [False, True]) +def test_empty_result_keeps_score_schema(docs, batch): + actual = query(docs, batch, pre_filter="id < 0").select(["id"]).with_score().order_by_score().to_arrow() + for result in actual if batch else [actual]: + assert result.num_rows == 0 + assert result.schema == pa.schema([("id", pa.int64()), ("_score", pa.float64())]) + + +def test_batch_keeps_one_lookup_and_query_specific_scores(docs): + original = ScanQuery._read_global_index_result + calls = [] + + def lookup(query, result): + calls.append(True) + return original(query, result) + + with patch.object(ScanQuery, "_read_global_index_result", lookup): + actual = query(docs, True).select(["id"]).with_score().order_by_score().to_arrow() + assert calls == [True] + assert actual[0].to_pylist()[0] == {"id": 1, "_score": 1.} + assert actual[1].to_pylist()[0] == {"id": 0, "_score": 1.} + + +@pytest.mark.parametrize("selector", ["snapshot", "tag"]) +def test_metadata_only_keeps_historical_deletions(docs, selector): + snapshot = docs.raw_table.snapshot_manager().get_latest_snapshot() + docs.raw_table.create_tag("saved", snapshot_id=snapshot.id) + options = {"snapshot_id": snapshot.id} if selector == "snapshot" else {"tag_name": "saved"} + before = query(docs, **options).select(["_ROW_ID"]).with_score().order_by_score() + expected = before.to_list() + docs.delete("id = 1") + assert before.to_list() == expected + assert query(docs).select(["_ROW_ID"]).with_score().to_arrow().num_rows == 3 + + +def test_commit_between_search_and_lookup_keeps_one_snapshot(docs): + search = query(docs).select(["id", "text"]).with_score().order_by_score() + expected = search.to_list() + original = ScanQuery._read_global_index_result + + def lookup(execution, result): + docs.update("id = 1", {"text": "changed"}) + return original(execution, result) + + with patch.object(ScanQuery, "_read_global_index_result", lookup): + assert search.to_list() == expected + assert search.to_list()[0]["text"] == "changed" + + +@pytest.mark.parametrize("name", [None, "", 1, "id", "_ROW_ID"]) +def test_score_alias_validation(docs, name): + with pytest.raises(ValueError, match="Score column name"): + query(docs).with_score(name) + + +def test_scores_require_data_evolution(docs): + docs.raw_table = docs.raw_table.copy({"data-evolution.enabled": "false"}) + with pytest.raises(NotImplementedError, match="data-evolution"): + query(docs).with_score() + with pytest.raises(NotImplementedError, match="data-evolution"): + query(docs).order_by_score() + + +def test_full_text_and_hybrid_scores(docs): + pytest.importorskip("paimon_ftindex") + docs.raw_table.copy({"deletion-vectors.enabled": "false"}).create_global_index("text", "full-text") + for search in [docs.search("paimon", column="text"), docs.search_hybrid([ + pm.vector_route("embedding", [0., 1.]), pm.text_route("paimon", column="text")])]: + result = search.select(["id"]).with_score().order_by_score().limit(3).to_arrow() + assert result.num_rows == 3 + assert result["_score"].to_pylist() == sorted(result["_score"].to_pylist(), reverse=True) + assert all(score > 0 for score in result["_score"].to_pylist())