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([