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 fd6a65e83ec4..51159f32167f 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_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 @@ def _rerank_indexed(self, plan, candidates, index_type): 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 @@ def reranked_iter(): 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 fe9d96536719..5dc7cc0650d0 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 @@ def __init__( 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 @@ def _open_offset_reader(self, vector_index_files, row_range_start, row_range_end 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 @@ def _read_raw_search(self, raw_row_ranges, pre_filter, query_vector, 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): @@ -444,8 +473,7 @@ def _maybe_rerank_indexed_results(self, results, index_type, query_vectors, 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(), @@ -488,6 +516,7 @@ def __init__(self, table, limit, vector_column, query_vector, filter_=None, 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() @@ -550,6 +579,7 @@ def __init__(self, table, limit, vector_column, query_vectors, 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: @@ -759,41 +789,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 376468a31cc5..afa55e15ed3d 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] + + +@pytest.mark.parametrize("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 000000000000..8b761efadab7 --- /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()} + + +@unittest.skipUnless(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 4d0c39e29eae..09c7c1f3034e 100644 --- a/paimon-python/pypaimon/tests/vector_search_filter_test.py +++ b/paimon-python/pypaimon/tests/vector_search_filter_test.py @@ -3510,6 +3510,9 @@ def test_batch_refine_factor_reranks_each_query(self): 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([