Skip to content

Commit a856ce9

Browse files
authored
Merge pull request #2079 from weaviate/trengrj/hybrid-diversity
Add hybrid diversity support
2 parents 904ea78 + 7c8c0e7 commit a856ce9

10 files changed

Lines changed: 305 additions & 0 deletions

File tree

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,149 @@
1+
"""Integration tests for hybrid search + MMR diversity selection.
2+
3+
``DiversitySelection`` passed inside ``HybridVector.near_vector`` /
4+
``HybridVector.near_text`` is applied by the server as a post-fusion MMR pass
5+
(Weaviate >= 1.39.0). These tests assert that ``balance=0`` (pure diversity)
6+
produces a different ordering than ``balance=1`` (pure relevance), and that
7+
``mmr.limit`` caps the result count.
8+
9+
The equivalent ``near_vector`` behaviour is covered in
10+
``test_collection_diversity.py``.
11+
"""
12+
13+
import pytest
14+
15+
from integration.conftest import CollectionFactory
16+
from weaviate.classes.query import Diversity, HybridVector
17+
from weaviate.collections.classes.config import Configure, DataType, Property
18+
from weaviate.collections.classes.data import DataObject
19+
20+
MIN_VERSION = (1, 38, 6)
21+
22+
23+
def _skip_if_unsupported(collection) -> None:
24+
if collection._connection._weaviate_version.is_lower_than(*MIN_VERSION):
25+
pytest.skip("Hybrid diversity selection requires Weaviate >= 1.38.6")
26+
27+
28+
def _create_clustered_collection(collection_factory: CollectionFactory):
29+
"""Create a collection with 3 tight clusters (a, b, c) of vectors in 3D."""
30+
collection = collection_factory(
31+
properties=[Property(name="text", data_type=DataType.TEXT)],
32+
vectorizer_config=Configure.Vectorizer.none(),
33+
)
34+
_skip_if_unsupported(collection)
35+
collection.data.insert_many(
36+
[
37+
DataObject(properties={"text": "a1"}, vector=[1.0, 0.0, 0.0]),
38+
DataObject(properties={"text": "a2"}, vector=[0.95, 0.05, 0.0]),
39+
DataObject(properties={"text": "a3"}, vector=[0.9, 0.1, 0.0]),
40+
DataObject(properties={"text": "b1"}, vector=[0.0, 1.0, 0.0]),
41+
DataObject(properties={"text": "b2"}, vector=[0.05, 0.95, 0.0]),
42+
DataObject(properties={"text": "c1"}, vector=[0.0, 0.0, 1.0]),
43+
]
44+
)
45+
return collection
46+
47+
48+
def _create_large_collection(collection_factory: CollectionFactory, n_items: int = 50):
49+
"""Create a collection with enough items (>25) that a small mmr.limit is distinguishable from the server's default limit."""
50+
collection = collection_factory(
51+
properties=[Property(name="text", data_type=DataType.TEXT)],
52+
vectorizer_config=Configure.Vectorizer.none(),
53+
)
54+
_skip_if_unsupported(collection)
55+
collection.data.insert_many(
56+
[
57+
DataObject(properties={"text": f"t{i}"}, vector=[1.0 - 0.001 * i, 0.0, 0.0])
58+
for i in range(n_items)
59+
]
60+
)
61+
return collection
62+
63+
64+
def test_hybrid_near_vector_balance_0_differs_from_balance_1(
65+
collection_factory: CollectionFactory,
66+
) -> None:
67+
"""Hybrid near-vector: balance=0 (diversity) must reorder vs balance=1 (relevance)."""
68+
collection = _create_clustered_collection(collection_factory)
69+
balance_0 = collection.query.hybrid(
70+
query=None,
71+
vector=HybridVector.near_vector(
72+
vector=[1.0, 0.0, 0.0],
73+
diversity_selection=Diversity.mmr(limit=3, balance=0.0),
74+
),
75+
limit=3,
76+
).objects
77+
balance_1 = collection.query.hybrid(
78+
query=None,
79+
vector=HybridVector.near_vector(
80+
vector=[1.0, 0.0, 0.0],
81+
diversity_selection=Diversity.mmr(limit=3, balance=1.0),
82+
),
83+
limit=3,
84+
).objects
85+
assert [o.uuid for o in balance_0] != [o.uuid for o in balance_1]
86+
87+
88+
def test_hybrid_near_vector_balance_1_matches_baseline(
89+
collection_factory: CollectionFactory,
90+
) -> None:
91+
"""Hybrid near-vector with MMR balance=1 (pure relevance) matches the plain baseline."""
92+
collection = _create_clustered_collection(collection_factory)
93+
baseline = collection.query.hybrid(
94+
query=None,
95+
vector=HybridVector.near_vector(vector=[1.0, 0.0, 0.0]),
96+
limit=3,
97+
).objects
98+
mmr_balance_1 = collection.query.hybrid(
99+
query=None,
100+
vector=HybridVector.near_vector(
101+
vector=[1.0, 0.0, 0.0],
102+
diversity_selection=Diversity.mmr(limit=3, balance=1.0),
103+
),
104+
limit=3,
105+
).objects
106+
assert [o.uuid for o in baseline] == [o.uuid for o in mmr_balance_1]
107+
108+
109+
def test_hybrid_alpha_1_balance_0_differs_from_balance_1(
110+
collection_factory: CollectionFactory,
111+
) -> None:
112+
"""Hybrid with explicit alpha=1.0 (pure vector) applies MMR like near_vector."""
113+
collection = _create_clustered_collection(collection_factory)
114+
balance_0 = collection.query.hybrid(
115+
query="irrelevant",
116+
alpha=1.0,
117+
vector=HybridVector.near_vector(
118+
vector=[1.0, 0.0, 0.0],
119+
diversity_selection=Diversity.mmr(limit=3, balance=0.0),
120+
),
121+
limit=3,
122+
).objects
123+
balance_1 = collection.query.hybrid(
124+
query="irrelevant",
125+
alpha=1.0,
126+
vector=HybridVector.near_vector(
127+
vector=[1.0, 0.0, 0.0],
128+
diversity_selection=Diversity.mmr(limit=3, balance=1.0),
129+
),
130+
limit=3,
131+
).objects
132+
assert [o.uuid for o in balance_0] != [o.uuid for o in balance_1]
133+
134+
135+
def test_hybrid_respects_mmr_limit(
136+
collection_factory: CollectionFactory,
137+
) -> None:
138+
"""Hybrid respects mmr.limit as the result-count cap when no outer limit is set."""
139+
mmr_limit = 5
140+
collection = _create_large_collection(collection_factory, n_items=50)
141+
142+
result = collection.query.hybrid(
143+
query=None,
144+
vector=HybridVector.near_vector(
145+
vector=[1.0, 0.0, 0.0],
146+
diversity_selection=Diversity.mmr(limit=mmr_limit, balance=0.5),
147+
),
148+
).objects
149+
assert len(result) == mmr_limit
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
"""Unit tests: hybrid search wires diversity_selection into the gRPC request.
2+
3+
Hybrid diversity is a post-fusion, hybrid-level operation, so the top-level
4+
``query.hybrid`` / ``generate.hybrid`` ``diversity_selection`` argument must
5+
populate the top-level ``Hybrid.selection.mmr`` in the SearchRequest proto (not
6+
the nested ``near_vector`` / ``near_text`` selection).
7+
"""
8+
9+
from weaviate.collections.grpc.query import _QueryGRPC
10+
from weaviate.classes.query import Diversity, HybridVector
11+
from weaviate.util import _ServerVersion
12+
13+
14+
_DEFAULT_VERSION = _ServerVersion(1, 38, 0)
15+
16+
17+
def _builder(version: _ServerVersion = _DEFAULT_VERSION) -> _QueryGRPC:
18+
return _QueryGRPC(
19+
weaviate_version=version,
20+
name="Dummy",
21+
tenant=None,
22+
consistency_level=None,
23+
validate_arguments=True,
24+
uses_125_api=True,
25+
uses_127_api=True,
26+
)
27+
28+
29+
def test_hybrid_near_vector_sets_top_level_selection() -> None:
30+
req = _builder().hybrid(
31+
query=None,
32+
vector=HybridVector.near_vector(vector=[1.0, 0.0, 0.0]),
33+
diversity_selection=Diversity.mmr(limit=7, balance=0.0),
34+
limit=7,
35+
)
36+
# Canonical location: top-level Hybrid.selection, not the nested near_vector.
37+
mmr = req.hybrid_search.selection.mmr
38+
assert mmr.limit == 7
39+
assert mmr.balance == 0.0
40+
assert not req.hybrid_search.near_vector.HasField("selection")
41+
42+
43+
def test_hybrid_near_text_sets_top_level_selection() -> None:
44+
req = _builder().hybrid(
45+
query=None,
46+
vector=HybridVector.near_text(query="cats"),
47+
diversity_selection=Diversity.mmr(limit=3, balance=0.5),
48+
limit=3,
49+
)
50+
mmr = req.hybrid_search.selection.mmr
51+
assert mmr.limit == 3
52+
assert mmr.balance == 0.5
53+
assert not req.hybrid_search.near_text.HasField("selection")
54+
55+
56+
def test_hybrid_without_selection_leaves_it_unset() -> None:
57+
req = _builder().hybrid(
58+
query=None,
59+
vector=HybridVector.near_vector(vector=[1.0, 0.0, 0.0]),
60+
limit=5,
61+
)
62+
assert not req.hybrid_search.HasField("selection")

weaviate/collections/grpc/query.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -173,6 +173,7 @@ def hybrid(
173173
generative: Optional[_Generative] = None,
174174
rerank: Optional[Rerank] = None,
175175
boost: Optional[_Boost] = None,
176+
diversity_selection: Optional[MMR] = None,
176177
target_vector: Optional[TargetVectorJoinType] = None,
177178
) -> search_get_pb2.SearchRequest:
178179
return self.__create_request(
@@ -196,6 +197,7 @@ def hybrid(
196197
fusion_type,
197198
distance,
198199
target_vector,
200+
diversity_selection,
199201
),
200202
)
201203

weaviate/collections/grpc/shared.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -600,6 +600,7 @@ def _parse_hybrid(
600600
fusion_type: Optional[HybridFusion],
601601
distance: Optional[NUMBER],
602602
target_vector: Optional[TargetVectorJoinType],
603+
diversity_selection: Optional[MMR] = None,
603604
) -> Union[base_search_pb2.Hybrid, None]:
604605
if self._validate_arguments:
605606
_validate_input(
@@ -739,6 +740,7 @@ def _parse_hybrid(
739740
vector_bytes=vector_bytes,
740741
vector_distance=distance,
741742
vectors=vectors,
743+
selection=self._diversity_selection_to_grpc(diversity_selection),
742744
bm25_search_operator=base_search_pb2.SearchOperatorOptions(
743745
operator=bm25_operator.operator,
744746
minimum_or_tokens_match=bm25_operator.minimum_should_match

weaviate/collections/queries/hybrid/generate/async_.pyi

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ from typing import Generic, List, Literal, Optional, Type, Union, overload
33
from weaviate.collections.classes.filters import FilterReturn
44
from weaviate.collections.classes.grpc import (
55
METADATA,
6+
MMR,
67
PROPERTIES,
78
REFERENCES,
89
BM25OperatorOptions,
@@ -56,6 +57,7 @@ class _HybridGenerateAsync(
5657
group_by: Literal[None] = None,
5758
rerank: Optional[Rerank] = None,
5859
boost: Optional[_Boost] = None,
60+
diversity_selection: Optional[MMR] = None,
5961
target_vector: Optional[TargetVectorJoinType] = None,
6062
include_vector: INCLUDE_VECTOR = False,
6163
return_metadata: Optional[METADATA] = None,
@@ -84,6 +86,7 @@ class _HybridGenerateAsync(
8486
group_by: Literal[None] = None,
8587
rerank: Optional[Rerank] = None,
8688
boost: Optional[_Boost] = None,
89+
diversity_selection: Optional[MMR] = None,
8790
target_vector: Optional[TargetVectorJoinType] = None,
8891
include_vector: INCLUDE_VECTOR = False,
8992
return_metadata: Optional[METADATA] = None,
@@ -112,6 +115,7 @@ class _HybridGenerateAsync(
112115
group_by: Literal[None] = None,
113116
rerank: Optional[Rerank] = None,
114117
boost: Optional[_Boost] = None,
118+
diversity_selection: Optional[MMR] = None,
115119
target_vector: Optional[TargetVectorJoinType] = None,
116120
include_vector: INCLUDE_VECTOR = False,
117121
return_metadata: Optional[METADATA] = None,
@@ -140,6 +144,7 @@ class _HybridGenerateAsync(
140144
group_by: Literal[None] = None,
141145
rerank: Optional[Rerank] = None,
142146
boost: Optional[_Boost] = None,
147+
diversity_selection: Optional[MMR] = None,
143148
target_vector: Optional[TargetVectorJoinType] = None,
144149
include_vector: INCLUDE_VECTOR = False,
145150
return_metadata: Optional[METADATA] = None,
@@ -168,6 +173,7 @@ class _HybridGenerateAsync(
168173
group_by: Literal[None] = None,
169174
rerank: Optional[Rerank] = None,
170175
boost: Optional[_Boost] = None,
176+
diversity_selection: Optional[MMR] = None,
171177
target_vector: Optional[TargetVectorJoinType] = None,
172178
include_vector: INCLUDE_VECTOR = False,
173179
return_metadata: Optional[METADATA] = None,
@@ -196,6 +202,7 @@ class _HybridGenerateAsync(
196202
group_by: Literal[None] = None,
197203
rerank: Optional[Rerank] = None,
198204
boost: Optional[_Boost] = None,
205+
diversity_selection: Optional[MMR] = None,
199206
target_vector: Optional[TargetVectorJoinType] = None,
200207
include_vector: INCLUDE_VECTOR = False,
201208
return_metadata: Optional[METADATA] = None,
@@ -224,6 +231,7 @@ class _HybridGenerateAsync(
224231
group_by: GroupBy,
225232
rerank: Optional[Rerank] = None,
226233
boost: Optional[_Boost] = None,
234+
diversity_selection: Optional[MMR] = None,
227235
target_vector: Optional[TargetVectorJoinType] = None,
228236
include_vector: INCLUDE_VECTOR = False,
229237
return_metadata: Optional[METADATA] = None,
@@ -252,6 +260,7 @@ class _HybridGenerateAsync(
252260
group_by: GroupBy,
253261
rerank: Optional[Rerank] = None,
254262
boost: Optional[_Boost] = None,
263+
diversity_selection: Optional[MMR] = None,
255264
target_vector: Optional[TargetVectorJoinType] = None,
256265
include_vector: INCLUDE_VECTOR = False,
257266
return_metadata: Optional[METADATA] = None,
@@ -280,6 +289,7 @@ class _HybridGenerateAsync(
280289
group_by: GroupBy,
281290
rerank: Optional[Rerank] = None,
282291
boost: Optional[_Boost] = None,
292+
diversity_selection: Optional[MMR] = None,
283293
target_vector: Optional[TargetVectorJoinType] = None,
284294
include_vector: INCLUDE_VECTOR = False,
285295
return_metadata: Optional[METADATA] = None,
@@ -308,6 +318,7 @@ class _HybridGenerateAsync(
308318
group_by: GroupBy,
309319
rerank: Optional[Rerank] = None,
310320
boost: Optional[_Boost] = None,
321+
diversity_selection: Optional[MMR] = None,
311322
target_vector: Optional[TargetVectorJoinType] = None,
312323
include_vector: INCLUDE_VECTOR = False,
313324
return_metadata: Optional[METADATA] = None,
@@ -336,6 +347,7 @@ class _HybridGenerateAsync(
336347
group_by: GroupBy,
337348
rerank: Optional[Rerank] = None,
338349
boost: Optional[_Boost] = None,
350+
diversity_selection: Optional[MMR] = None,
339351
target_vector: Optional[TargetVectorJoinType] = None,
340352
include_vector: INCLUDE_VECTOR = False,
341353
return_metadata: Optional[METADATA] = None,
@@ -364,6 +376,7 @@ class _HybridGenerateAsync(
364376
group_by: GroupBy,
365377
rerank: Optional[Rerank] = None,
366378
boost: Optional[_Boost] = None,
379+
diversity_selection: Optional[MMR] = None,
367380
target_vector: Optional[TargetVectorJoinType] = None,
368381
include_vector: INCLUDE_VECTOR = False,
369382
return_metadata: Optional[METADATA] = None,
@@ -392,6 +405,7 @@ class _HybridGenerateAsync(
392405
group_by: Optional[GroupBy] = None,
393406
rerank: Optional[Rerank] = None,
394407
boost: Optional[_Boost] = None,
408+
diversity_selection: Optional[MMR] = None,
395409
target_vector: Optional[TargetVectorJoinType] = None,
396410
include_vector: INCLUDE_VECTOR = False,
397411
return_metadata: Optional[METADATA] = None,

0 commit comments

Comments
 (0)