Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
93 changes: 57 additions & 36 deletions paimon-python/pypaimon/table/source/vector_search_read.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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"

Comment thread
TheR1sing3un marked this conversation as resolved.
inferred = None
for key, value in list(options.items()) + list(table_map.items()):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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)
Loading
Loading