diff --git a/python/tests/test_collection_fts.py b/python/tests/test_collection_fts.py index 6cc884b..fab6446 100644 --- a/python/tests/test_collection_fts.py +++ b/python/tests/test_collection_fts.py @@ -172,7 +172,9 @@ class TestFtsOnlyCollectionLifecycle: class TestFtsOnlyCollectionQueryValidation: def test_vector_query_rejected(self, fts_collection: Collection): """Vector query on a no-vector collection must raise.""" - with pytest.raises(ValueError, match="No vector field found"): + with pytest.raises( + ValueError, match="Vector field 'content' not found in schema" + ): fts_collection.query( queries=Query(field_name="content", vector=[0.1, 0.2, 0.3]), topk=5, @@ -181,7 +183,9 @@ class TestFtsOnlyCollectionQueryValidation: def test_id_query_rejected(self, fts_collection: Collection): """ID-based query on a no-vector collection must raise.""" fts_collection.insert(_make_docs()[:1]) - with pytest.raises(ValueError, match="No vector field found"): + with pytest.raises( + ValueError, match="Vector field 'content' not found in schema" + ): fts_collection.query( queries=Query(field_name="content", id="pk_0"), topk=5, diff --git a/python/tests/test_collection_group_by.py b/python/tests/test_collection_group_by.py new file mode 100644 index 0000000..fbb4add --- /dev/null +++ b/python/tests/test_collection_group_by.py @@ -0,0 +1,348 @@ +# Copyright 2025-present the zvec project +# +# Licensed 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 __future__ import annotations + +import pytest +import zvec +from zvec import ( + Collection, + CollectionOption, + DataType, + Doc, + FieldSchema, + FlatIndexParam, + Fts, + GroupResult, + HnswIndexParam, + HnswQueryParam, + InvertIndexParam, + Query, + VectorSchema, +) + +# ==================== Constants ==================== + +GB_DIMENSION = 4 +GB_NUM_DOCS = 12 +GB_NUM_GROUPS = 3 +GB_TOPK_PER_GROUP = 2 + + +# ==================== Fixtures ==================== + + +@pytest.fixture(scope="session") +def group_by_collection_schema(): + """Collection schema for group-by end-to-end tests. + + Mirrors the data layout in ``vector_column_indexer_test.cc``: + a 4-dimensional dense vector and a scalar ``group_id`` field used + for grouping. + """ + return zvec.CollectionSchema( + name="test_group_by_collection", + fields=[ + FieldSchema( + "id", + DataType.INT64, + nullable=False, + index_param=InvertIndexParam(enable_range_optimization=True), + ), + FieldSchema( + "group_id", + DataType.INT64, + nullable=False, + index_param=InvertIndexParam(), + ), + ], + vectors=[ + VectorSchema( + "dense", + DataType.VECTOR_FP32, + dimension=GB_DIMENSION, + index_param=HnswIndexParam(), + ), + VectorSchema( + "dense_flat", + DataType.VECTOR_FP32, + dimension=GB_DIMENSION, + index_param=FlatIndexParam(), + ), + ], + ) + + +@pytest.fixture(scope="session") +def collection_option(): + return CollectionOption(read_only=False, enable_mmap=True) + + +@pytest.fixture +def group_by_docs(): + """Generate docs matching the C++ GroupByIndexerTest fixture. + + Doc ``i`` has vector ``[i, i, i, i]`` and ``group_id = i % 3``. + """ + return [ + Doc( + id=f"{i}", + fields={"id": i, "group_id": i % GB_NUM_GROUPS}, + vectors={ + "dense": [float(i)] * GB_DIMENSION, + "dense_flat": [float(i)] * GB_DIMENSION, + }, + ) + for i in range(GB_NUM_DOCS) + ] + + +@pytest.fixture(scope="function") +def group_by_collection( + tmp_path_factory, group_by_collection_schema, collection_option +) -> Collection: + """Function-scoped fixture: creates and opens a collection for group-by tests.""" + temp_dir = tmp_path_factory.mktemp("zvec_group_by") + collection_path = temp_dir / "test_group_by_collection" + + coll = zvec.create_and_open( + path=str(collection_path), + schema=group_by_collection_schema, + option=collection_option, + ) + + assert coll is not None, "Failed to create and open group-by collection" + assert coll.path == str(collection_path) + assert coll.schema.name == group_by_collection_schema.name + + try: + yield coll + finally: + if hasattr(coll, "destroy") and coll is not None: + try: + coll.destroy() + except Exception as e: + print(f"Warning: failed to destroy collection: {e}") + + +@pytest.fixture +def group_by_collection_with_docs( + group_by_collection: Collection, group_by_docs +) -> Collection: + """Setup: insert group-by fixture docs.""" + assert group_by_collection.stats.doc_count == 0 + result = group_by_collection.insert(group_by_docs) + assert len(result) == len(group_by_docs) + for item in result: + assert item.ok() + assert group_by_collection.stats.doc_count == len(group_by_docs) + + yield group_by_collection + + # Teardown + group_by_collection.delete([doc.id for doc in group_by_docs]) + + +# ==================== Helpers ==================== + + +def _assert_grouped_results(results, num_groups, topk_per_group, query_value): + """Validate group-by result structure and ordering. + + Each returned group must: + - contain only docs whose ``group_id`` matches ``group_by_value`` + - have at most ``topk_per_group`` docs + - have docs sorted by descending score + """ + assert len(results) == num_groups, ( + f"Expected {num_groups} groups, got {len(results)}" + ) + + group_values = set() + for group in results: + assert isinstance(group, GroupResult) + group_value = int(group.group_by_value) + group_values.add(group_value) + docs = group.docs + assert 1 <= len(docs) <= topk_per_group + + for doc in docs: + assert int(doc.field("group_id")) == group_value + + scores = [doc.score for doc in docs] + assert scores == sorted(scores, reverse=True), ( + "Docs must be sorted by score desc" + ) + + # Score sanity: for query [1,1,1,1] and vector [i,i,i,i], + # IP score is 4 * i. + for doc in docs: + doc_id = int(doc.field("id")) + expected_score = float(doc_id * sum(query_value)) + assert abs(doc.score - expected_score) < 0.1 + + assert group_values == set(range(num_groups)) + + +# ==================== Tests ==================== + + +@pytest.mark.usefixtures("group_by_collection_with_docs") +class TestGroupBySearch: + def test_group_by_defaults(self, group_by_collection: Collection): + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=[1.0] * GB_DIMENSION), + group_by_field_name="group_id", + ) + assert len(results) == 2 + assert all(1 <= len(group.docs) <= 3 for group in results) + + def test_group_by_hnsw(self, group_by_collection: Collection): + """Group-by search over an HNSW index.""" + query_vector = [1.0] * GB_DIMENSION + results = group_by_collection.group_by_query( + Query( + field_name="dense", vector=query_vector, param=HnswQueryParam(ef=300) + ), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + ) + _assert_grouped_results(results, GB_NUM_GROUPS, GB_TOPK_PER_GROUP, query_vector) + + def test_group_by_flat(self, group_by_collection: Collection): + """Group-by search over a FLAT index.""" + query_vector = [1.0] * GB_DIMENSION + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=query_vector), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + ) + _assert_grouped_results(results, GB_NUM_GROUPS, GB_TOPK_PER_GROUP, query_vector) + + def test_group_by_with_filter(self, group_by_collection: Collection): + """Group-by search with a scalar filter.""" + query_vector = [1.0] * GB_DIMENSION + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=query_vector), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + filter="id < 6", + ) + # Only docs 0..5 are visible; every group still has at least one doc. + assert len(results) == GB_NUM_GROUPS + for group in results: + for doc in group.docs: + assert int(doc.field("id")) < 6 + + def test_group_by_include_vector(self, group_by_collection: Collection): + """Group-by search returns original vectors when requested.""" + query_vector = [1.0] * GB_DIMENSION + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=query_vector), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + include_vector=True, + ) + assert len(results) == GB_NUM_GROUPS + for group in results: + for doc in group.docs: + vec = doc.vector("dense_flat") + doc_id = int(doc.field("id")) + assert vec == pytest.approx([float(doc_id)] * GB_DIMENSION, abs=1e-5) + + def test_group_by_output_fields(self, group_by_collection: Collection): + """Group-by search honors scalar output field selection.""" + query_vector = [1.0] * GB_DIMENSION + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=query_vector), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + output_fields=["group_id"], + ) + assert len(results) == GB_NUM_GROUPS + for group in results: + for doc in group.docs: + assert doc.has_field("group_id") + + def test_group_by_invalid_field(self, group_by_collection: Collection): + """Group-by with a non-existent vector field raises an error.""" + with pytest.raises(ValueError): + group_by_collection.group_by_query( + Query(field_name="nonexistent", vector=[1.0] * GB_DIMENSION), + group_by_field_name="group_id", + ) + + def test_group_by_query_by_id(self, group_by_collection: Collection): + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", id="11"), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + ) + assert len(results) == GB_NUM_GROUPS + + @pytest.mark.parametrize( + ("query", "error"), + [ + (Query(field_name="content", fts=Fts(match_string="text")), "FTS"), + (Query(field_name="dense_flat"), "vector or document id"), + ], + ) + def test_group_by_rejects_unsupported_query( + self, group_by_collection: Collection, query: Query, error: str + ): + with pytest.raises(ValueError, match=error): + group_by_collection.group_by_query(query, "group_id") + + @pytest.mark.parametrize( + ("kwargs", "error"), + [ + ({"group_by_field_name": ""}, "group_by_field_name"), + ({"group_by_field_name": "group_id", "group_count": 0}, "group_count"), + ( + {"group_by_field_name": "group_id", "topk_per_group": 0}, + "topk_per_group", + ), + ], + ) + def test_group_by_rejects_invalid_group_params( + self, group_by_collection: Collection, kwargs: dict, error: str + ): + with pytest.raises(ValueError, match=error): + group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=[1.0] * GB_DIMENSION), + **kwargs, + ) + + +class TestGroupByEmptyCollection: + def test_group_by_empty_collection(self, group_by_collection: Collection): + """Group-by on an empty collection returns an empty list.""" + results = group_by_collection.group_by_query( + Query(field_name="dense_flat", vector=[1.0] * GB_DIMENSION), + group_by_field_name="group_id", + group_count=GB_NUM_GROUPS, + topk_per_group=GB_TOPK_PER_GROUP, + ) + assert results == [] + + +def test_group_by_public_api_exports(): + assert zvec.GroupResult is GroupResult + assert not hasattr(zvec, "GroupByQuery") + assert not hasattr(Collection, "groupby_query") diff --git a/python/zvec/__init__.py b/python/zvec/__init__.py index d75a391..e06d23f 100644 --- a/python/zvec/__init__.py +++ b/python/zvec/__init__.py @@ -95,7 +95,7 @@ from .model import schema as schema # —— Core data structures —— from .model.collection import Collection -from .model.doc import Doc, DocList +from .model.doc import Doc, DocList, GroupResult # —— Query & index parameters —— # —— FTS params (C++ binding) —— @@ -161,6 +161,7 @@ __all__ = [ "VectorSchema", "CollectionStats", # Parameters + "GroupResult", "Query", "VectorQuery", "Fts", diff --git a/python/zvec/__init__.pyi b/python/zvec/__init__.pyi index 3e75f93..adc23d9 100644 --- a/python/zvec/__init__.pyi +++ b/python/zvec/__init__.pyi @@ -11,7 +11,7 @@ from .extension import ReRanker, RrfReRanker, WeightedReRanker from .extension.embedding import DenseEmbeddingFunction from .model import param, schema from .model.collection import Collection -from .model.doc import Doc, DocList +from .model.doc import Doc, DocList, GroupResult from .model.param import ( AddColumnOption, AlterColumnOption, @@ -66,6 +66,7 @@ __all__: list = [ "Fts", "FtsIndexParam", "FtsQueryParam", + "GroupResult", "HnswIndexParam", "HnswQueryParam", "HnswRabitqIndexParam", @@ -132,7 +133,7 @@ class _Collection: include_vector: bool = True, ) -> dict[str, _Doc]: ... def Flush(self) -> None: ... - def GroupByQuery(self, arg0: ...) -> list[...]: ... + def GroupByQuery(self, arg0: param._GroupByVectorQuery) -> list[_GroupResult]: ... def Insert(self, arg0: collections.abc.Sequence[_Doc]) -> list[typing.Status]: ... def Optimize(self, arg0: param.OptimizeOption) -> None: ... def Options(self) -> param.CollectionOption: ... @@ -165,6 +166,12 @@ class _Doc: def set_pk(self, arg0: str) -> None: ... def set_score(self, arg0: typing.SupportsFloat) -> None: ... +class _GroupResult: + @property + def docs(self) -> list[_Doc]: ... + @property + def group_by_value(self) -> str: ... + class _DocOp: """ Members: diff --git a/python/zvec/executor/query_executor.py b/python/zvec/executor/query_executor.py index 3faf160..cadbb17 100644 --- a/python/zvec/executor/query_executor.py +++ b/python/zvec/executor/query_executor.py @@ -223,6 +223,31 @@ class QueryExecutor: fts.match_string = query.fts.match_string or "" search_query.fts = fts + def set_query_vector( + self, query: Query, search_query, collection: _Collection + ) -> None: + """Resolve a Query vector and set it on a native query object.""" + vector_schema = self._schema.vector(query.field_name) + if vector_schema is None: + raise ValueError(f"Vector field '{query.field_name}' not found in schema") + + if query.has_vector(): + vec_data = query.vector + elif query.has_id(): + fetched = collection.Fetch([query.id]) + doc = next(iter(fetched.values()), None) + if not doc: + raise ValueError(f"Document with id '{query.id}' not found") + vec_data = doc.get_any(vector_schema.name, vector_schema.data_type) + else: + raise ValueError("Group by query requires a vector or document id") + + target_dtype = DTYPE_MAP.get(vector_schema.data_type.value) + search_query.set_vector( + vector_schema._get_object(), + convert_to_numpy(vec_data, target_dtype) if target_dtype else vec_data, + ) + def _build_search_query( self, ctx: QueryContext, query: Query, collection: _Collection ) -> _SearchQuery: @@ -235,32 +260,6 @@ class QueryExecutor: # set FTS query if provided self._apply_fts(query, search_query) - vector_schema = None if query.has_vector() or query.has_id(): - vector_schema = ( - self._schema.vector(query.field_name) - if query - else self._schema.vectors[0] - ) - - if vector_schema is None: - raise ValueError("No vector field found") - - # set vector - if query.has_vector(): - vec_data = query.vector - elif query.has_id(): - fetched = collection.Fetch([query.id]) - doc = next(iter(fetched.values()), None) - if not doc: - raise ValueError(f"Document with id '{query.id}' not found") - vec_data = doc.get_any(vector_schema.name, vector_schema.data_type) - else: - return search_query - - target_dtype = DTYPE_MAP.get(vector_schema.data_type.value) - search_query.set_vector( - vector_schema._get_object(), - convert_to_numpy(vec_data, target_dtype) if target_dtype else vec_data, - ) + self.set_query_vector(query, search_query, collection) return search_query diff --git a/python/zvec/model/__init__.py b/python/zvec/model/__init__.py index 7d5b068..2a5d882 100644 --- a/python/zvec/model/__init__.py +++ b/python/zvec/model/__init__.py @@ -14,7 +14,7 @@ from __future__ import annotations from .collection import Collection -from .doc import Doc +from .doc import Doc, GroupResult from .param.query import Fts, Query, VectorQuery from .schema.collection_schema import CollectionSchema from .schema.field_schema import FieldSchema @@ -25,6 +25,7 @@ __all__ = [ "Doc", "FieldSchema", "Fts", + "GroupResult", "Query", "VectorQuery", ] diff --git a/python/zvec/model/collection.py b/python/zvec/model/collection.py index b90caae..6663a1b 100644 --- a/python/zvec/model/collection.py +++ b/python/zvec/model/collection.py @@ -17,12 +17,13 @@ import warnings from typing import Optional, Union, overload from zvec._zvec import _Collection +from zvec._zvec.param import _GroupByVectorQuery from ..executor import QueryContext, QueryExecutor from ..extension import ReRanker from ..typing import Status from .convert import convert_to_cpp_doc, convert_to_py_doc -from .doc import Doc, DocList +from .doc import Doc, DocList, GroupResult from .param import ( AddColumnOption, AlterColumnOption, @@ -42,6 +43,11 @@ from .schema import CollectionSchema, CollectionStats, FieldSchema __all__ = ["Collection"] +def _require_positive_integer(value, name: str) -> None: + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer") + + class Collection: """Represents an opened collection in Zvec. @@ -437,3 +443,80 @@ class Collection: reranker=reranker, ) return self._querier.execute(ctx, self._obj) + + def group_by_query( + self, + query: Query, + group_by_field_name: str, + group_count=2, + topk_per_group=3, + *, + filter: Optional[str] = None, + include_vector: bool = False, + output_fields: Optional[list[str]] = None, + ) -> list[GroupResult]: + """Perform group-by vector search. + + Groups results by a scalar field value, returning the top-k documents + within each group, ordered by similarity score. + + Accepts a single vector Query. Full-text search is not supported. + + Args: + query (Query): Vector query. + group_by_field_name (str): Scalar field used to group results. + group_count (int): Maximum number of groups to return. + topk_per_group (int): Maximum number of documents in each group. + filter (Optional[str]): Boolean expression used to filter candidates. + include_vector (bool): Whether returned documents include vectors. + output_fields (Optional[list[str]]): Scalar fields to return. + + Returns: + list[GroupResult]: Grouped documents sorted by score. + + Examples: + >>> results = collection.group_by_query( + ... zvec.Query( + ... field_name="embedding", + ... vector=[0.1, 0.2, 0.3], + ... param=zvec.HnswQueryParam(ef=300), + ... ), + ... group_by_field_name="category", + ... group_count=5, + ... topk_per_group=3, + ... ) + >>> for group in results: + ... print(group.group_by_value, len(group.docs)) + """ + query._validate() + if query.has_fts(): + raise ValueError("Group by query does not support FTS") + if not query.has_vector() and not query.has_id(): + raise ValueError("Group by query requires a vector or document id") + if not group_by_field_name: + raise ValueError("group_by_field_name cannot be empty") + _require_positive_integer(group_count, "group_count") + _require_positive_integer(topk_per_group, "topk_per_group") + cpp_query = _GroupByVectorQuery() + cpp_query.field_name = query.field_name + cpp_query.group_by_field_name = group_by_field_name + cpp_query.group_count = group_count + cpp_query.topk_per_group = topk_per_group + cpp_query.include_vector = include_vector + if filter: + cpp_query.filter = filter + if output_fields is not None: + cpp_query.output_fields = output_fields + if query.param: + cpp_query.query_params = query.param + self._querier.set_query_vector(query, cpp_query, self._obj) + + raw_results = self._obj.GroupByQuery(cpp_query) + + return [ + GroupResult( + group_by_value=group.group_by_value, + docs=[convert_to_py_doc(doc, self.schema) for doc in group.docs], + ) + for group in raw_results + ] diff --git a/python/zvec/model/doc.py b/python/zvec/model/doc.py index 175c946..83f75cb 100644 --- a/python/zvec/model/doc.py +++ b/python/zvec/model/doc.py @@ -14,6 +14,7 @@ from __future__ import annotations import json +from dataclasses import dataclass from typing import Any, Optional from ..common import VectorType @@ -21,6 +22,7 @@ from ..common import VectorType __all__ = [ "Doc", "DocList", + "GroupResult", ] @@ -176,3 +178,16 @@ class Doc: #: Type alias for query results: a list of documents returned by a single query route. DocList = list[Doc] + + +@dataclass(frozen=True) +class GroupResult: + """A group value and its matching documents. + + Attributes: + group_by_value (str): String representation of the grouped scalar field value. + docs (list[Doc]): Matching documents in similarity order. + """ + + group_by_value: str + docs: list[Doc] diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 0c50b32..947468f 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -5771,7 +5771,7 @@ zvec_group_by_vector_query_t *zvec_group_by_vector_query_create(void) { ZVEC_TRY_RETURN_NULL( "Failed to create GroupByVectorQuery", auto *query = new zvec::GroupByVectorQuery(); - query->group_count_ = 2; query->group_topk_ = 3; + query->group_count_ = 2; query->topk_per_group_ = 3; return reinterpret_cast(query);) return nullptr; } @@ -5841,23 +5841,23 @@ uint32_t zvec_group_by_vector_query_get_group_count( return ptr->group_count_; } -zvec_error_code_t zvec_group_by_vector_query_set_group_topk( - zvec_group_by_vector_query_t *query, uint32_t topk) { +zvec_error_code_t zvec_group_by_vector_query_set_topk_per_group( + zvec_group_by_vector_query_t *query, uint32_t topk_per_group) { if (!query) { SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "GroupByVectorQuery pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } auto *ptr = reinterpret_cast(query); - ptr->group_topk_ = topk; + ptr->topk_per_group_ = topk_per_group; return ZVEC_OK; } -uint32_t zvec_group_by_vector_query_get_group_topk( +uint32_t zvec_group_by_vector_query_get_topk_per_group( const zvec_group_by_vector_query_t *query) { if (!query) return 3; auto *ptr = reinterpret_cast(query); - return ptr->group_topk_; + return ptr->topk_per_group_; } zvec_error_code_t zvec_group_by_vector_query_set_query_vector( diff --git a/src/binding/python/model/param/python_param.cc b/src/binding/python/model/param/python_param.cc index 500455c..3c3ff62 100644 --- a/src/binding/python/model/param/python_param.cc +++ b/src/binding/python/model/param/python_param.cc @@ -1952,6 +1952,71 @@ Args: })); } +void set_query_vector(QueryTarget &target, const FieldSchema &field_schema, + const py::object &obj) { + const DataType data_type = field_schema.data_type(); + + // Dense vector data is referenced by the query object. Callers + // must not modify the source data until the query returns. + if (FieldSchema::is_dense_vector_field(data_type)) { + if (!py::isinstance(obj)) { + throw py::type_error("Dense vector[" + field_schema.name() + + "] expects a ndarray, got " + + std::string(py::str(py::type::of(obj)))); + } + const auto arr = obj.cast(); + if (arr.ndim() != 1) { + throw py::type_error("Dense vector expects 1D array, got " + + std::to_string(arr.ndim()) + "D"); + } + const auto buf = arr.request(); + target.clause_ = VectorViewClause{ + std::string_view(static_cast(buf.ptr), + static_cast(buf.size) * buf.itemsize), + {}, + {}}; + return; + } + // sparse vector + if (FieldSchema::is_sparse_vector_field(data_type)) { + if (!py::isinstance(obj)) { + throw py::type_error("Sparse vector[" + field_schema.name() + + "] expects a Python dict, got " + + std::string(py::str(py::type::of(obj)))); + } + const auto sparse = obj.cast(); + + switch (data_type) { + case DataType::SPARSE_VECTOR_FP16: { + auto [indices, values] = serialize_sparse_vector( + sparse, [](const py::handle &h, size_t idx) { + float f = checked_cast( + h, "Sparse value[" + std::to_string(idx) + "]", "FLOAT"); + return ailego::Float16(f); + }); + target.set_sparse_vector(std::move(indices), std::move(values)); + break; + } + case DataType::SPARSE_VECTOR_FP32: { + auto [indices, values] = serialize_sparse_vector( + sparse, [](const py::handle &h, size_t idx) { + return checked_cast( + h, "Sparse value[" + std::to_string(idx) + "]", "FLOAT"); + }); + target.set_sparse_vector(std::move(indices), std::move(values)); + break; + } + default: + throw py::type_error("Unsupported sparse vector type: " + + std::to_string(static_cast(data_type))); + } + return; + } + + throw py::type_error("Unsupported vector field type for field: " + + field_schema.name()); +} + void ZVecPyParams::bind_vector_query(py::module_ &m) { // bind Fts py::class_(m, "_Fts") @@ -2038,74 +2103,7 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { "set_vector", [](SearchQuery &self, const FieldSchema &field_schema, const py::object &obj) { - const DataType data_type = field_schema.data_type(); - - // Dense vector data is referenced by the query object. Callers - // must not modify the source data until the query returns. - if (FieldSchema::is_dense_vector_field(data_type)) { - if (!py::isinstance(obj)) { - throw py::type_error("Dense vector[" + field_schema.name() + - "] expects a ndarray, got " + - std::string(py::str(py::type::of(obj)))); - } - const auto arr = obj.cast(); - if (arr.ndim() != 1) { - throw py::type_error("Dense vector expects 1D array, got " + - std::to_string(arr.ndim()) + "D"); - } - const auto buf = arr.request(); - self.target_.clause_ = VectorViewClause{ - std::string_view( - static_cast(buf.ptr), - static_cast(buf.size) * buf.itemsize), - {}, - {}}; - return; - } - // sparse vector - if (FieldSchema::is_sparse_vector_field(data_type)) { - if (!py::isinstance(obj)) { - throw py::type_error("Sparse vector[" + field_schema.name() + - "] expects a Python dict, got " + - std::string(py::str(py::type::of(obj)))); - } - const auto sparse = obj.cast(); - - switch (data_type) { - case DataType::SPARSE_VECTOR_FP16: { - auto [indices, values] = - serialize_sparse_vector( - sparse, [](const py::handle &h, size_t idx) { - float f = checked_cast( - h, "Sparse value[" + std::to_string(idx) + "]", - "FLOAT"); - return ailego::Float16(f); - }); - self.target_.set_sparse_vector(std::move(indices), - std::move(values)); - break; - } - case DataType::SPARSE_VECTOR_FP32: { - auto [indices, values] = serialize_sparse_vector( - sparse, [](const py::handle &h, size_t idx) { - return checked_cast( - h, "Sparse value[" + std::to_string(idx) + "]", - "FLOAT"); - }); - self.target_.set_sparse_vector(std::move(indices), - std::move(values)); - break; - } - default: - throw py::type_error( - "Unsupported sparse vector type: " + - std::to_string(static_cast(data_type))); - } - return; - } - - throw py::type_error("Unsupported vector field type for field: " + - field_schema.name()); + set_query_vector(self.target_, field_schema, obj); }, py::arg("field_schema"), py::arg("obj"), py::keep_alive<1, 3>(), "Set query vector. Dense vector source data must not be modified " @@ -2265,5 +2263,37 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { } return obj; })); + + // _GroupByVectorQuery + py::class_(m, "_GroupByVectorQuery") + .def(py::init<>()) + .def_property( + "field_name", + [](const GroupByVectorQuery &q) { return q.target_.field_name_; }, + [](GroupByVectorQuery &q, std::string v) { + q.target_.field_name_ = std::move(v); + }) + .def_readwrite("filter", &GroupByVectorQuery::filter_) + .def_readwrite("include_vector", &GroupByVectorQuery::include_vector_) + .def_readwrite("output_fields", &GroupByVectorQuery::output_fields_) + .def_readwrite("group_by_field_name", + &GroupByVectorQuery::group_by_field_name_) + .def_readwrite("group_count", &GroupByVectorQuery::group_count_) + .def_readwrite("topk_per_group", &GroupByVectorQuery::topk_per_group_) + .def_property( + "query_params", + [](const GroupByVectorQuery &q) { return q.target_.query_params_; }, + [](GroupByVectorQuery &q, QueryParams::Ptr p) { + q.target_.query_params_ = std::move(p); + }) + .def( + "set_vector", + [](GroupByVectorQuery &self, const FieldSchema &field_schema, + const py::object &obj) { + set_query_vector(self.target_, field_schema, obj); + }, + py::arg("field_schema"), py::arg("obj"), py::keep_alive<1, 3>(), + "Set query vector. Dense vector source data must not be modified " + "until the query finishes."); } } // namespace zvec diff --git a/src/binding/python/model/python_collection.cc b/src/binding/python/model/python_collection.cc index b1311f1..4468a53 100644 --- a/src/binding/python/model/python_collection.cc +++ b/src/binding/python/model/python_collection.cc @@ -48,6 +48,10 @@ T unwrap_expected(const tl::expected &exp) { } void ZVecPyCollection::Initialize(pybind11::module_ &m) { + py::class_(m, "_GroupResult") + .def_readonly("group_by_value", &GroupResult::group_by_value_) + .def_readonly("docs", &GroupResult::docs_); + py::class_ collection(m, "_Collection"); bind_db_methods(collection); bind_ddl_methods(collection); @@ -281,7 +285,6 @@ void ZVecPyCollection::bind_dql_methods( py::gil_scoped_release release; result = self.GroupByQuery(query); } - // return GroupResults return unwrap_expected(result); }) .def( @@ -313,4 +316,4 @@ void ZVecPyCollection::bind_dql_methods( "for introspection and testing only; not part of the stable API."); } -} // namespace zvec \ No newline at end of file +} // namespace zvec diff --git a/src/core/algorithm/hnsw/hnsw_context.h b/src/core/algorithm/hnsw/hnsw_context.h index b7ac79c..4e2f90c 100644 --- a/src/core/algorithm/hnsw/hnsw_context.h +++ b/src/core/algorithm/hnsw/hnsw_context.h @@ -50,8 +50,8 @@ class HnswContext : public IndexContext { public: //! Set topk of search result void set_topk(uint32_t val) override { - topk_ = val; - topk_heap_.limit(std::max(val, ef_)); + topk_ = group_by_search() ? group_topk_ * group_num_ : val; + topk_heap_.limit(std::max(topk_, ef_)); } //! Retrieve search result @@ -530,12 +530,9 @@ class HnswContext : public IndexContext { void set_group_params(uint32_t group_num, uint32_t group_topk) override { group_num_ = group_num; group_topk_ = group_topk; - - topk_ = group_topk_ * group_num_; - - topk_heap_.limit(std::max(topk_, ef_)); - group_topk_heaps_.clear(); + + set_topk(group_topk_ * group_num_); } private: diff --git a/src/core/algorithm/hnsw_rabitq/hnsw_rabitq_context.h b/src/core/algorithm/hnsw_rabitq/hnsw_rabitq_context.h index ea2c86b..9321886 100644 --- a/src/core/algorithm/hnsw_rabitq/hnsw_rabitq_context.h +++ b/src/core/algorithm/hnsw_rabitq/hnsw_rabitq_context.h @@ -49,8 +49,8 @@ class HnswRabitqContext : public IndexContext { public: //! Set topk of search result void set_topk(uint32_t val) override { - topk_ = val; - topk_heap_.limit(std::max(val, ef_)); + topk_ = group_by_search() ? group_topk_ * group_num_ : val; + topk_heap_.limit(std::max(topk_, ef_)); } //! Retrieve search result @@ -470,12 +470,9 @@ class HnswRabitqContext : public IndexContext { void set_group_params(uint32_t group_num, uint32_t group_topk) override { group_num_ = group_num; group_topk_ = group_topk; - - topk_ = group_topk_ * group_num_; - - topk_heap_.limit(std::max(topk_, ef_)); - group_topk_heaps_.clear(); + + set_topk(group_topk_ * group_num_); } void set_provider(IndexProvider::Pointer provider) { diff --git a/src/core/algorithm/hnsw_sparse/hnsw_sparse_context.h b/src/core/algorithm/hnsw_sparse/hnsw_sparse_context.h index 56ae1cb..07778d1 100644 --- a/src/core/algorithm/hnsw_sparse/hnsw_sparse_context.h +++ b/src/core/algorithm/hnsw_sparse/hnsw_sparse_context.h @@ -43,8 +43,8 @@ class HnswSparseContext : public IndexContext { public: //! Set topk of search result void set_topk(uint32_t val) override { - topk_ = val; - topk_heap_.limit(std::max(val, ef_)); + topk_ = group_by_search() ? group_topk_ * group_num_ : val; + topk_heap_.limit(std::max(topk_, ef_)); } //! Retrieve search result @@ -481,12 +481,9 @@ class HnswSparseContext : public IndexContext { void set_group_params(uint32_t group_num, uint32_t group_topk) override { group_num_ = group_num; group_topk_ = group_topk; - - topk_ = group_topk_ * group_num_; - - topk_heap_.limit(std::max(topk_, ef_)); - group_topk_heaps_.clear(); + + set_topk(group_topk_ * group_num_); } private: diff --git a/src/core/interface/index.cc b/src/core/interface/index.cc index 6c69530..755f673 100644 --- a/src/core/interface/index.cc +++ b/src/core/interface/index.cc @@ -1030,6 +1030,11 @@ int Index::_get_coarse_search_topk( return floor(search_param->topk * scale_factor); } +// Set or clear group-by state on a pooled context before each search. +// +// Set group-by state before topk so contexts can derive the effective candidate +// count from the active group parameters. This order also ensures pooled +// contexts restore the ordinary topk after stale group-by state is cleared. void Index::_set_group_by_on_context( const BaseIndexQueryParam::Pointer &search_param, core::IndexContext::Pointer &context) { @@ -1037,6 +1042,9 @@ void Index::_set_group_by_on_context( context->set_group_by(search_param->group_by_param->group_by); context->set_group_params(search_param->group_by_param->group_count, search_param->group_by_param->group_topk); + } else { + context->set_group_params(0, 0); + context->reset_group_by(); } } diff --git a/src/core/interface/indexes/hnsw_index.cc b/src/core/interface/indexes/hnsw_index.cc index 9756289..7b88263 100644 --- a/src/core/interface/indexes/hnsw_index.cc +++ b/src/core/interface/indexes/hnsw_index.cc @@ -149,6 +149,9 @@ int HNSWIndex::_prepare_for_search( return core::IndexError_Runtime; } + // Set group state first so set_topk() derives the effective candidate count. + _set_group_by_on_context(search_param, context); + context->set_topk(hnsw_search_param->topk); context->set_fetch_vector(hnsw_search_param->fetch_vector); if (hnsw_search_param->filter) { @@ -168,7 +171,7 @@ int HNSWIndex::_prepare_for_search( std::min(256u, hnsw_search_param->prefetch_lines); params.set(core::PARAM_HNSW_STREAMER_PL, real_search_pl); context->update(params); - _set_group_by_on_context(search_param, context); + return 0; } @@ -182,4 +185,4 @@ int HNSWIndex::_get_coarse_search_topk( return ret; } -} // namespace zvec::core_interface \ No newline at end of file +} // namespace zvec::core_interface diff --git a/src/core/interface/indexes/hnsw_rabitq_index.cc b/src/core/interface/indexes/hnsw_rabitq_index.cc index c555389..251fe10 100644 --- a/src/core/interface/indexes/hnsw_rabitq_index.cc +++ b/src/core/interface/indexes/hnsw_rabitq_index.cc @@ -108,6 +108,9 @@ int HNSWRabitqIndex::_prepare_for_search( return core::IndexError_Runtime; } + // Set group state first so set_topk() derives the effective candidate count. + _set_group_by_on_context(search_param, context); + context->set_topk(hnsw_search_param->topk); context->set_fetch_vector(hnsw_search_param->fetch_vector); if (hnsw_search_param->filter) { @@ -121,7 +124,7 @@ int HNSWRabitqIndex::_prepare_for_search( std::max(1u, std::min(2048u, hnsw_search_param->ef_search)); params.set(core::PARAM_HNSW_RABITQ_STREAMER_EF, real_search_ef); context->update(params); - _set_group_by_on_context(search_param, context); + return 0; #endif // RABITQ_SUPPORTED } diff --git a/src/db/sqlengine/sqlengine_impl.cc b/src/db/sqlengine/sqlengine_impl.cc index c216646..8a84bb5 100644 --- a/src/db/sqlengine/sqlengine_impl.cc +++ b/src/db/sqlengine/sqlengine_impl.cc @@ -121,7 +121,7 @@ Result SQLEngineImpl::execute_group_by( auto first_query_info = build_query_info( collection, query, std::make_shared(group_by_query.group_by_field_name_, - group_by_query.group_topk_, + group_by_query.topk_per_group_, group_by_query.group_count_)); if (!first_query_info) { return tl::make_unexpected(first_query_info.error()); @@ -144,7 +144,7 @@ Result SQLEngineImpl::execute_group_by( auto query_info = build_query_info( collection, query, std::make_shared(group_by_query.group_by_field_name_, - group_by_query.group_topk_, + group_by_query.topk_per_group_, group_by_query.group_count_)); if (!query_info) { return tl::make_unexpected(query_info.error()); diff --git a/src/include/zvec/c_api.h b/src/include/zvec/c_api.h index 855b420..169dad6 100644 --- a/src/include/zvec/c_api.h +++ b/src/include/zvec/c_api.h @@ -2149,21 +2149,21 @@ ZVEC_EXPORT uint32_t ZVEC_CALL zvec_group_by_vector_query_get_group_count( const zvec_group_by_vector_query_t *query); /** - * @brief Set group topk + * @brief Set the maximum number of results per group * @param query Group by vector query pointer - * @param topk Number of results per group + * @param topk_per_group Number of results per group * @return zvec_error_code_t Error code */ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL -zvec_group_by_vector_query_set_group_topk(zvec_group_by_vector_query_t *query, - uint32_t topk); +zvec_group_by_vector_query_set_topk_per_group( + zvec_group_by_vector_query_t *query, uint32_t topk_per_group); /** - * @brief Get group topk + * @brief Get the maximum number of results per group * @param query Group by vector query pointer * @return uint32_t Number of results per group */ -ZVEC_EXPORT uint32_t ZVEC_CALL zvec_group_by_vector_query_get_group_topk( +ZVEC_EXPORT uint32_t ZVEC_CALL zvec_group_by_vector_query_get_topk_per_group( const zvec_group_by_vector_query_t *query); /** diff --git a/src/include/zvec/core/framework/index_context.h b/src/include/zvec/core/framework/index_context.h index 1cf0e03..1414214 100644 --- a/src/include/zvec/core/framework/index_context.h +++ b/src/include/zvec/core/framework/index_context.h @@ -19,7 +19,7 @@ #include #include #include -#include +#include #include #include diff --git a/src/include/zvec/core/framework/index_groupby.h b/src/include/zvec/core/framework/index_group_by.h similarity index 100% rename from src/include/zvec/core/framework/index_groupby.h rename to src/include/zvec/core/framework/index_group_by.h diff --git a/src/include/zvec/core/interface/index.h b/src/include/zvec/core/interface/index.h index 6c01d70..118d06d 100644 --- a/src/include/zvec/core/interface/index.h +++ b/src/include/zvec/core/interface/index.h @@ -205,7 +205,8 @@ class Index { const BaseIndexQueryParam::Pointer &search_param); //! Helper: set group_by on context from the query param (common for all - //! index types). Call this at the end of _prepare_for_search. + //! index types). Call this before set_topk() when topk depends on group + //! state. static void _set_group_by_on_context( const BaseIndexQueryParam::Pointer &search_param, core::IndexContext::Pointer &context); diff --git a/src/include/zvec/db/query.h b/src/include/zvec/db/query.h index 6de901d..6d30ccc 100644 --- a/src/include/zvec/db/query.h +++ b/src/include/zvec/db/query.h @@ -171,7 +171,7 @@ struct GroupByVectorQuery { std::optional> output_fields_; std::string group_by_field_name_; uint32_t group_count_{2}; - uint32_t group_topk_{3}; + uint32_t topk_per_group_{3}; }; struct GroupResult { diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index cbd7488..7a5a4f7 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -6315,6 +6315,13 @@ void test_diskann_wiring_on_vector_query(void) { zvec_group_by_vector_query_t *gbq = zvec_group_by_vector_query_create(); TEST_ASSERT(gbq != NULL); + TEST_ASSERT(zvec_group_by_vector_query_get_topk_per_group(gbq) == 3); + err = zvec_group_by_vector_query_set_topk_per_group(gbq, 7); + TEST_ASSERT(err == ZVEC_OK); + TEST_ASSERT(zvec_group_by_vector_query_get_topk_per_group(gbq) == 7); + err = zvec_group_by_vector_query_set_topk_per_group(NULL, 7); + TEST_ASSERT(err == ZVEC_ERROR_INVALID_ARGUMENT); + zvec_diskann_query_params_t *dp2 = zvec_query_params_diskann_create(200); TEST_ASSERT(dp2 != NULL); err = zvec_group_by_vector_query_set_diskann_params(gbq, dp2); diff --git a/tests/db/sqlengine/sqlengine_test.cc b/tests/db/sqlengine/sqlengine_test.cc index 9f4894b..9001235 100644 --- a/tests/db/sqlengine/sqlengine_test.cc +++ b/tests/db/sqlengine/sqlengine_test.cc @@ -146,7 +146,7 @@ TEST_F(SqlEngineTest, GroupBy) { GroupByVectorQuery query; query.group_by_field_name_ = "name"; query.group_count_ = 3; - query.group_topk_ = 2; + query.topk_per_group_ = 2; query.output_fields_ = {"id", "name", "score"}; query.filter_ = "id > 3 and score < 0.1"; if (const char *env_var = std::getenv("FILTER"); env_var != nullptr) {