This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git


The following commit(s) were added to refs/heads/master by this push:
     new 92666078c3 [python] Use persisted vector metrics for raw search and 
refinement (#9754)
92666078c3 is described below

commit 92666078c34a91c5a233b4e701d38582d034a4bf
Author: chaoyang <[email protected]>
AuthorDate: Sun Sep 13 16:26:26 2026 +0800

    [python] Use persisted vector metrics for raw search and refinement (#9754)
---
 .../table/source/primary_key_vector_read.py        |   9 +-
 .../pypaimon/table/source/vector_search_read.py    |  93 +++++----
 .../tests/primary_key_global_index_golden_test.py  |  46 +++++
 .../tests/vector_metric_consistency_test.py        | 210 +++++++++++++++++++++
 .../pypaimon/tests/vector_search_filter_test.py    |   3 +
 5 files changed, 320 insertions(+), 41 deletions(-)

diff --git a/paimon-python/pypaimon/table/source/primary_key_vector_read.py 
b/paimon-python/pypaimon/table/source/primary_key_vector_read.py
index fd6a65e83e..51159f3216 100644
--- a/paimon-python/pypaimon/table/source/primary_key_vector_read.py
+++ b/paimon-python/pypaimon/table/source/primary_key_vector_read.py
@@ -25,7 +25,7 @@ from pypaimon.table.source.primary_key_scored_result import (
 from pypaimon.table.source.primary_key_vector_scan import 
PrimaryKeyVectorScanPlan
 from pypaimon.table.source.vector_search_read import DataEvolutionVectorRead
 from pypaimon.table.source.vector_search_read import (
-    _check_vector_dimension, _compute_score, _raw_search_metric, 
_to_vector_list)
+    _check_vector_dimension, _compute_score, _to_vector_list)
 from pypaimon.read.split import DataSplit
 from pypaimon.globalindex.indexed_split import IndexedSplit
 from pypaimon.deletionvectors.deletion_vector import DeletionVector
@@ -38,6 +38,7 @@ class PrimaryKeyVectorRead(DataEvolutionVectorRead):
     def read_plan(self, plan):
         if not isinstance(plan, PrimaryKeyVectorScanPlan):
             raise ValueError("Primary-key vector read requires a 
PrimaryKeyVectorScanPlan.")
+        self._index_metric = None
         index_type = self._table.options.primary_key_vector_index_type(
             self._vector_column.name)
         indexed_limit = self._indexed_search_limit(index_type)
@@ -115,8 +116,7 @@ class PrimaryKeyVectorRead(DataEvolutionVectorRead):
             plan.snapshot_id, source_splits, candidates)
         reader = self._table.new_read_builder().with_projection(
             [self._vector_column.name]).new_read()
-        metric = _raw_search_metric(
-            self._table, self._vector_column, self._options, index_type)
+        metric = self._search_metric(index_type)
 
         def reranked_iter():
             for split in candidate_result.splits:
@@ -168,8 +168,7 @@ class PrimaryKeyVectorRead(DataEvolutionVectorRead):
         return reranked
 
     def _raw_candidates(self, plan):
-        metric = _raw_search_metric(
-            self._table, self._vector_column, self._options,
+        metric = self._search_metric(
             self._table.options.primary_key_vector_index_type(
                 self._vector_column.name))
         read_builder = self._table.new_read_builder().with_projection(
diff --git a/paimon-python/pypaimon/table/source/vector_search_read.py 
b/paimon-python/pypaimon/table/source/vector_search_read.py
index 3e86150333..5f5dc6ec70 100644
--- a/paimon-python/pypaimon/table/source/vector_search_read.py
+++ b/paimon-python/pypaimon/table/source/vector_search_read.py
@@ -86,6 +86,31 @@ class AbstractVectorSearchReadImpl:
         self._filter = filter_
         self._partition_filter = partition_filter
         self._options = dict(options or {})
+        self._index_metric = None
+
+    def _search_metric(self, index_type=None):
+        if self._index_metric is not None:
+            return self._index_metric
+        return _raw_search_metric(
+            self._table, self._vector_column, self._options, index_type)
+
+    def _record_index_metric(self, reader, index_type):
+        """Keep one persisted metric for indexed scores, raw search and 
refinement."""
+        metric_getter = getattr(reader, "vector_metric", None)
+        if metric_getter is None:
+            return
+        metric = _normalize_metric(metric_getter())
+        requested = _configured_vector_metric(
+            self._options, self._vector_column, index_type)
+        if requested is not None and requested != metric:
+            raise ValueError(
+                "Query vector metric '%s' does not match index metric '%s' for 
column '%s'."
+                % (requested, metric, self._vector_column.name))
+        if self._index_metric is not None and self._index_metric != metric:
+            raise ValueError(
+                "Cannot merge vector indexes with different metrics '%s' and 
'%s' for column '%s'."
+                % (self._index_metric, metric, self._vector_column.name))
+        self._index_metric = metric
 
     def _pre_filters(self, splits, snapshot=None):
         # type: (list) -> List[RoaringBitmap64]
@@ -235,7 +260,12 @@ class AbstractVectorSearchReadImpl:
             index_io_meta_list,
             self._table.table_schema.options,
         )
-        return reader, OffsetGlobalIndexReader(reader, row_range_start, 
row_range_end)
+        try:
+            self._record_index_metric(reader, vector_index_files[0].index_type)
+            return reader, OffsetGlobalIndexReader(reader, row_range_start, 
row_range_end)
+        except Exception:
+            reader.close()
+            raise
 
     def _eval(self, row_range_start, row_range_end, vector_index_files,
               query_vector, search_limit, include_row_ids):
@@ -275,8 +305,7 @@ class AbstractVectorSearchReadImpl:
             return DictBasedScoredIndexResult({})
 
         top_k_heap = []
-        metric = _raw_search_metric(
-            self._table, self._vector_column, self._options, index_type)
+        metric = self._search_metric(index_type)
         row_ids = table.column(SpecialFields.ROW_ID.name).to_pylist()
         vectors = table.column(self._vector_column.name).to_pylist()
         for row_id, stored in zip(row_ids, vectors):
@@ -448,8 +477,7 @@ class AbstractVectorSearchReadImpl:
 
         raw_vectors = self._read_raw_vectors(
             union_candidates, include_filter=False, snapshot=snapshot)
-        metric = _raw_search_metric(
-            self._table, self._vector_column, self._options, index_type)
+        metric = self._search_metric(index_type)
         return [
             self._score_raw_vectors(
                 candidates[i].results(),
@@ -492,6 +520,7 @@ class DataEvolutionVectorRead(AbstractVectorSearchReadImpl, 
VectorSearchRead):
         self._query_vector = query_vector
 
     def _read(self, splits, snapshot):
+        self._index_metric = None
         index_splits, raw_splits = _split_search_splits(splits)
         if not index_splits and not raw_splits:
             return GlobalIndexResult.create_empty()
@@ -554,6 +583,7 @@ class 
BatchVectorSearchReadImpl(AbstractVectorSearchReadImpl,
         self._query_vectors = list(query_vectors)
 
     def _read_batch(self, splits, snapshot):
+        self._index_metric = None
         n = len(self._query_vectors)
         index_splits, raw_splits = _split_search_splits(splits)
         if not index_splits and not raw_splits:
@@ -812,41 +842,32 @@ def _table_options_map(table):
     return table_options.to_map() if table_options is not None else {}
 
 
-def _raw_search_metric(table, vector_column, options, index_type=None):
-    candidates = []
+def _configured_vector_metric(options, vector_column, index_type=None):
     field_prefix = "fields.%s." % vector_column.name
     index_prefix = "%s." % index_type if index_type else None
-    for key in [
-        field_prefix + "distance.metric",
-        field_prefix + "metric",
-        *(([
-            index_prefix + "distance.metric",
-            index_prefix + "metric",
-        ]) if index_prefix is not None else []),
-        "test.vector.metric",
-        "lumina.distance.metric",
-        "distance.metric",
-        "metric",
-    ]:
+    keys = [field_prefix + "pk-vector.distance.metric",
+            field_prefix + "distance.metric", field_prefix + "metric"]
+    if index_prefix is not None:
+        keys.extend([index_prefix + "distance.metric", index_prefix + 
"metric"])
+    keys.extend(["test.vector.metric", "lumina.distance.metric", 
"distance.metric", "metric"])
+    for key in keys:
         if key in options:
-            candidates.append(options[key])
+            return _normalize_metric(options[key])
+    return None
+
+
+def _raw_search_metric(table, vector_column, options, index_type=None):
+    from pypaimon.globalindex.vindex.vindex_vector_global_index_reader import 
VINDEX_IDENTIFIERS
+
     table_map = _table_options_map(table)
-    for key in [
-        field_prefix + "distance.metric",
-        field_prefix + "metric",
-        *(([
-            index_prefix + "distance.metric",
-            index_prefix + "metric",
-        ]) if index_prefix is not None else []),
-        "test.vector.metric",
-        "lumina.distance.metric",
-        "distance.metric",
-        "metric",
-    ]:
-        if key in table_map:
-            candidates.append(table_map[key])
-    if candidates:
-        return _normalize_metric(candidates[0])
+    for source in (options, table_map):
+        configured = _configured_vector_metric(source, vector_column, 
index_type)
+        if configured is not None:
+            return configured
+
+    # Before an index exists, use its writer's default, not another column's 
metric.
+    if index_type in VINDEX_IDENTIFIERS:
+        return "inner_product"
 
     inferred = None
     for key, value in list(options.items()) + list(table_map.items()):
diff --git 
a/paimon-python/pypaimon/tests/primary_key_global_index_golden_test.py 
b/paimon-python/pypaimon/tests/primary_key_global_index_golden_test.py
index 376468a31c..afa55e15ed 100644
--- a/paimon-python/pypaimon/tests/primary_key_global_index_golden_test.py
+++ b/paimon-python/pypaimon/tests/primary_key_global_index_golden_test.py
@@ -92,6 +92,25 @@ def test_java_primary_key_vector_index(catalog):
     assert filtered_rows.column("id").to_pylist() == [3]
 
 
+def test_java_primary_key_vector_refinement_uses_persisted_metric(catalog):
+    _require_native("paimon_vindex")
+    table = catalog.get_table("default.test_pk_vector_golden")
+
+    def search(read_table):
+        return (read_table.new_vector_search_builder()
+                .with_vector_column("embedding")
+                .with_query_vector([1.0, 0.0, 0.0, 0.0])
+                .with_limit(3)
+                .with_option("ivf.refine_factor", "2")
+                .execute_local())
+
+    expected = search(table).positions
+    assert expected
+    for metric in ("l2", "inner_product", "cosine"):
+        changed = table.copy({"fields.embedding.distance.metric": metric})
+        assert search(changed).positions == expected
+
+
 def test_java_primary_key_full_text_index(catalog):
     _require_native("paimon_ftindex")
     table = catalog.get_table("default.test_pk_full_text_golden")
@@ -102,3 +121,30 @@ def test_java_primary_key_full_text_index(catalog):
               .execute_local())
     rows = _read_search_result(table, result)
     assert sorted(rows.column("id").to_pylist()) == [1, 3]
+
+
[email protected]("metric,expected_id", [("l2", 2), ("cosine", 1), 
("inner_product", 1)])
+def test_java_primary_key_raw_only_uses_column_metric(catalog, metric, 
expected_id):
+    from dataclasses import replace
+    from pypaimon.table.source.primary_key_vector_scan import 
PrimaryKeyVectorScanPlan
+    from pypaimon.table.source.vector_search_read import _compute_score
+
+    table = catalog.get_table("default.test_pk_vector_golden")
+    # Remove legacy aliases so only the documented PK column option is 
available.
+    options = {key: None for key in table.options.options.to_map() if 
key.endswith(".metric")}
+    options.update({"fields.embedding.pk-vector.distance.metric": metric,
+                    "fields.other.pk-vector.distance.metric": "cosine",
+                    "vector-index.search-mode": "full"})
+    table = table.copy(options)
+    query = [0.1, 0.0, 0.0, 0.0]
+    builder = 
(table.new_vector_search_builder().with_vector_column("embedding")
+               .with_query_vector(query).with_limit(1))
+    plan = builder.new_vector_search_scan().scan()
+    raw_plan = PrimaryKeyVectorScanPlan(plan.snapshot_id, [
+        replace(split, payloads=(), uncovered_data_files=tuple(
+            file.file_name for file in split.data_split.files)) for split in 
plan.splits()])
+    result = builder.new_vector_search_read().read_plan(raw_plan)
+    rows = _read_search_result(table, result)
+    assert rows.column("id").to_pylist() == [expected_id]
+    assert result.positions[0].score == _compute_score(
+        query, rows.column("embedding")[0].as_py(), metric)
diff --git a/paimon-python/pypaimon/tests/vector_metric_consistency_test.py 
b/paimon-python/pypaimon/tests/vector_metric_consistency_test.py
new file mode 100644
index 0000000000..8b761efada
--- /dev/null
+++ b/paimon-python/pypaimon/tests/vector_metric_consistency_test.py
@@ -0,0 +1,210 @@
+# 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.
+
+import importlib.util
+import unittest
+from unittest.mock import Mock, patch
+
+import pyarrow as pa
+
+from pypaimon.table.source.vector_search_read import (
+    BatchVectorSearchReadImpl, DataEvolutionVectorRead, _raw_search_metric)
+from pypaimon.tests.vector_search_filter_test import (
+    _StubTable, _entry, _field, _install_raw_vector_read_builder)
+from pypaimon.table.source.vector_search_split import RawVectorSearchSplit
+from pypaimon.utils.range import Range
+from pypaimon.tests.data_evolution_test_helpers import BatchModeMixin, 
DataEvolutionTestBase
+
+
+def _scores(result):
+    getter = result.score_getter()
+    return {row_id: getter(row_id) for row_id in result.results()}
+
+
[email protected](importlib.util.find_spec("paimon_vindex"), "paimon-vindex 
is not installed")
+class NativeVectorMetricTest(BatchModeMixin, DataEvolutionTestBase, 
unittest.TestCase):
+    pa_schema = pa.schema([('embedding', pa.list_(pa.float32()))])
+    table_options = {
+        'row-tracking.enabled': 'true', 'data-evolution.enabled': 'true',
+        'global-index.enabled': 'true', 'bucket': '-1', 'file.format': 
'parquet',
+        'vector-index.search-mode': 'full',
+    }
+
+    def _append(self, table, vectors):
+        self._write_arrow(table, pa.table({'embedding': vectors}, 
schema=self.pa_schema))
+
+    def _build(self, table, metric=None):
+        options = {'ivf-flat.dimension': '2', 'ivf-flat.nlist': '1'}
+        if metric is not None:
+            options['ivf-flat.distance.metric'] = metric
+        self.assertEqual(1, table.create_global_index('embedding', 
index_type='ivf-flat', options=options))
+
+    def _builder(self, table, batch=False):
+        if batch:
+            builder = 
table.new_batch_vector_search_builder().with_query_vectors([[1, 0], [1, 0]])
+        else:
+            builder = table.new_vector_search_builder().with_query_vector([1, 
0])
+        return 
builder.with_vector_column('embedding').with_limit(2).with_option('ivf.nprobe', 
'1')
+
+    def _execute(self, builder, batch):
+        return builder.execute_batch_local() if batch else 
[builder.execute_local()]
+
+    def test_mixed_search_uses_persisted_metric(self):
+        for metric, expected in ((None, {0: 2.0, 1: 3.0}),
+                                 ('inner_product', {0: 2.0, 1: 3.0}),
+                                 ('cosine', {0: 1.0, 1: 1.0}),
+                                 ('l2', {0: 0.5, 1: 0.2})):
+            with self.subTest(metric=metric):
+                table = self._create_table()
+                self._append(table, [[2, 0]])
+                self._build(table, metric)
+                self._append(table, [[3, 0]])
+                for batch in (False, True):
+                    for result in self._execute(self._builder(table, batch), 
batch):
+                        actual = _scores(result)
+                        self.assertEqual(set(expected), set(actual))
+                        for row_id, score in expected.items():
+                            self.assertAlmostEqual(score, actual[row_id], 
places=6)
+
+    def test_refinement_uses_persisted_metric(self):
+        for metric, expected in ((None, {1: 3.0}), ('cosine', {0: 1.0}), 
('l2', {0: 0.5})):
+            with self.subTest(metric=metric):
+                table = self._create_table()
+                self._append(table, [[2, 0], [3, 0]])
+                self._build(table, metric)
+                for batch in (False, True):
+                    builder = self._builder(table, 
batch).with_limit(1).with_option('ivf.refine_factor', '2')
+                    for result in self._execute(builder, batch):
+                        self.assertEqual(expected, _scores(result))
+
+    def test_persisted_metric_overrides_changed_table_options(self):
+        table = self._create_table()
+        self._append(table, [[2, 0]])
+        self._build(table)
+        self._append(table, [[3, 0]])
+        changed = table.copy({'fields.embedding.distance.metric': 'l2'})
+        for batch in (False, True):
+            for result in self._execute(self._builder(changed, batch), batch):
+                self.assertEqual({0: 2.0, 1: 3.0}, _scores(result))
+
+    def test_incompatible_query_metric_is_rejected(self):
+        table = self._create_table()
+        self._append(table, [[2, 0]])
+        self._build(table)
+        for batch in (False, True):
+            builder = self._builder(table, batch).with_option('metric', 'l2')
+            with self.assertRaisesRegex(ValueError, "Query vector metric 
'l2'.*index metric 'inner_product'"):
+                self._execute(builder, batch)
+
+    def test_incompatible_shard_metrics_are_rejected(self):
+        table = self._create_table()
+        self._append(table, [[2, 0]])
+        self._build(table)
+        self._append(table, [[3, 0]])
+        self._build(table, 'cosine')
+        for batch in (False, True):
+            with self.assertRaisesRegex(ValueError, 'Cannot merge vector 
indexes with different metrics'):
+                self._execute(self._builder(table, batch), batch)
+
+    def test_default_metric_is_consistent_before_and_after_build(self):
+        table = self._create_table()
+        self._append(table, [[2, 0], [3, 0]])
+        for indexed in (False, True):
+            if indexed:
+                self._build(table)
+            for batch in (False, True):
+                builder = self._builder(table, 
batch).with_option('index-type', 'ivf-flat').with_limit(1)
+                for result in self._execute(builder, batch):
+                    self.assertEqual({1: 3.0}, _scores(result))
+
+
+class VectorMetricResolutionTest(unittest.TestCase):
+    def setUp(self):
+        self.column = _field(1, 'embedding', 'FLOAT')
+        self.table = _StubTable(fields=[self.column], entries=[])
+        self.table.table_schema.options = {}
+
+    def _reader(self, options=None, batch=False):
+        kwargs = dict(table=self.table, vector_column=self.column, limit=1, 
options=options)
+        if batch:
+            return BatchVectorSearchReadImpl(query_vectors=[[1.0]], **kwargs)
+        return DataEvolutionVectorRead(query_vector=[1.0], **kwargs)
+
+    def test_query_metric_aliases_are_validated(self):
+        for key in ('fields.embedding.pk-vector.distance.metric',
+                    'fields.embedding.distance.metric', 
'fields.embedding.metric',
+                    'ivf-flat.distance.metric', 'ivf-flat.metric', 
'distance.metric', 'metric'):
+            with self.subTest(key=key):
+                native = Mock(spec=['vector_metric'])
+                native.vector_metric.return_value = 'inner_product'
+                reader = self._reader({key: 'inner-product'})
+                reader._record_index_metric(native, 'ivf-flat')
+                self.assertEqual('inner_product', 
reader._search_metric('ivf-flat'))
+                reader = self._reader({key: 'l2'})
+                with self.assertRaisesRegex(ValueError, 'does not match index 
metric'):
+                    reader._record_index_metric(native, 'ivf-flat')
+
+    def test_other_columns_do_not_override_persisted_metric(self):
+        reader = self._reader({'fields.other.metric': 'l2'})
+        native = Mock(spec=['vector_metric'])
+        native.vector_metric.return_value = 'cosine'
+        reader._record_index_metric(native, 'ivf-flat')
+        self.assertEqual('cosine', reader._search_metric('ivf-flat'))
+        self.assertEqual('inner_product', _raw_search_metric(
+            self.table, self.column, {'fields.other.metric': 'l2'}, 
'ivf-flat'))
+
+    def test_raw_only_defaults_match_vindex_writers(self):
+        for kind in ('ivf-flat', 'ivf-pq', 'ivf-sq', 'ivf-rq', 'diskann'):
+            self.assertEqual('inner_product', _raw_search_metric(self.table, 
self.column, {}, kind))
+            self.assertEqual('cosine', _raw_search_metric(
+                self.table, self.column, {'metric': 'cosine'}, kind))
+        self.assertEqual('l2', _raw_search_metric(self.table, self.column, {}))
+        self.assertEqual('l2', _raw_search_metric(self.table, self.column, {}, 
'lumina'))
+
+    def test_reader_closes_on_metadata_and_metric_errors(self):
+        entry = _entry(None, field_id=1, index_type='ivf-flat', 
file_name='vectors.index',
+                       row_range_start=0, row_range_end=1)
+        for failure in ('metadata', 'query', 'shard'):
+            with self.subTest(failure=failure):
+                native = Mock(spec=['vector_metric', 'close'])
+                native.vector_metric.return_value = 'inner_product'
+                reader = self._reader({'metric': 'l2'} if failure == 'query' 
else {})
+                if failure == 'metadata':
+                    native.vector_metric.side_effect = RuntimeError('invalid 
index metadata')
+                if failure == 'shard':
+                    previous = Mock(spec=['vector_metric'])
+                    previous.vector_metric.return_value = 'cosine'
+                    reader._record_index_metric(previous, 'ivf-flat')
+                with 
patch('pypaimon.table.source.vector_search_read._create_vector_reader', 
return_value=native):
+                    with self.assertRaises((RuntimeError, ValueError)):
+                        reader._open_offset_reader([entry.index_file], 0, 1)
+                native.close.assert_called_once_with()
+
+    def test_metric_does_not_leak_between_read_calls(self):
+        _install_raw_vector_read_builder(self.table, 'embedding', {0: [2.0], 
1: [3.0]})
+        split = RawVectorSearchSplit([Range(0, 1)], [], 'ivf-flat')
+        for batch in (False, True):
+            reader = self._reader(batch=batch)
+            native = Mock(spec=['vector_metric'])
+            native.vector_metric.return_value = 'l2'
+            reader._record_index_metric(native, 'ivf-flat')
+            results = reader.read_batch([split]) if batch else 
[reader.read([split])]
+            self.assertEqual([{1: 3.0}], [_scores(r) for r in results])
+
+
+if __name__ == '__main__':
+    unittest.main()
diff --git a/paimon-python/pypaimon/tests/vector_search_filter_test.py 
b/paimon-python/pypaimon/tests/vector_search_filter_test.py
index c43804e6f0..39daa2bbd5 100644
--- a/paimon-python/pypaimon/tests/vector_search_filter_test.py
+++ b/paimon-python/pypaimon/tests/vector_search_filter_test.py
@@ -3517,6 +3517,9 @@ class BatchVectorSearchTest(unittest.TestCase):
         def _fake_create(index_type, file_io, index_path,
                          index_io_meta_list, options=None):
             class _FakeReader(GlobalIndexReader):
+                def vector_metric(self_inner):
+                    return "l2"
+
                 def visit_batch_vector_search(self_inner, bvs):
                     captured_limits.append(bvs.limit)
                     return _completed_future([

Reply via email to