Skip to content

Commit 48c78a8

Browse files
committed
Add top level Hybrid diversity
1 parent 120a1b0 commit 48c78a8

10 files changed

Lines changed: 107 additions & 34 deletions

File tree

test/collection/test_hybrid_diversity.py

Lines changed: 13 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,19 +1,22 @@
11
"""Unit tests: hybrid search wires diversity_selection into the gRPC request.
22
3-
Hybrid diversity is a post-fusion, hybrid-level operation, so the
4-
``HybridVector.near_vector`` / ``HybridVector.near_text`` ``diversity_selection``
5-
argument must populate the top-level ``Hybrid.selection.mmr`` in the
6-
SearchRequest proto (not the nested ``near_vector`` / ``near_text`` selection).
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).
77
"""
88

99
from weaviate.collections.grpc.query import _QueryGRPC
1010
from weaviate.classes.query import Diversity, HybridVector
1111
from weaviate.util import _ServerVersion
1212

1313

14-
def _builder() -> _QueryGRPC:
14+
_DEFAULT_VERSION = _ServerVersion(1, 38, 0)
15+
16+
17+
def _builder(version: _ServerVersion = _DEFAULT_VERSION) -> _QueryGRPC:
1518
return _QueryGRPC(
16-
weaviate_version=_ServerVersion(1, 39, 0),
19+
weaviate_version=version,
1720
name="Dummy",
1821
tenant=None,
1922
consistency_level=None,
@@ -26,10 +29,8 @@ def _builder() -> _QueryGRPC:
2629
def test_hybrid_near_vector_sets_top_level_selection() -> None:
2730
req = _builder().hybrid(
2831
query=None,
29-
vector=HybridVector.near_vector(
30-
vector=[1.0, 0.0, 0.0],
31-
diversity_selection=Diversity.mmr(limit=7, balance=0.0),
32-
),
32+
vector=HybridVector.near_vector(vector=[1.0, 0.0, 0.0]),
33+
diversity_selection=Diversity.mmr(limit=7, balance=0.0),
3334
limit=7,
3435
)
3536
# Canonical location: top-level Hybrid.selection, not the nested near_vector.
@@ -42,10 +43,8 @@ def test_hybrid_near_vector_sets_top_level_selection() -> None:
4243
def test_hybrid_near_text_sets_top_level_selection() -> None:
4344
req = _builder().hybrid(
4445
query=None,
45-
vector=HybridVector.near_text(
46-
query="cats",
47-
diversity_selection=Diversity.mmr(limit=3, balance=0.5),
48-
),
46+
vector=HybridVector.near_text(query="cats"),
47+
diversity_selection=Diversity.mmr(limit=3, balance=0.5),
4948
limit=3,
5049
)
5150
mmr = req.hybrid_search.selection.mmr

weaviate/collections/classes/grpc.py

Lines changed: 0 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -760,7 +760,6 @@ class _HybridNearBase(_WeaviateInput):
760760

761761
distance: Optional[float] = None
762762
certainty: Optional[float] = None
763-
diversity_selection: Optional[MMR] = None
764763

765764

766765
class _HybridNearText(_HybridNearBase):
@@ -773,20 +772,17 @@ class _HybridNearVector: # can't be a Pydantic model because of validation issu
773772
vector: NearVectorInputType
774773
distance: Optional[float]
775774
certainty: Optional[float]
776-
diversity_selection: Optional[MMR]
777775

778776
def __init__(
779777
self,
780778
*,
781779
vector: NearVectorInputType,
782780
distance: Optional[float] = None,
783781
certainty: Optional[float] = None,
784-
diversity_selection: Optional[MMR] = None,
785782
) -> None:
786783
self.vector = vector
787784
self.distance = distance
788785
self.certainty = certainty
789-
self.diversity_selection = diversity_selection
790786

791787

792788
HybridVectorType = Union[NearVectorInputType, _HybridNearText, _HybridNearVector]
@@ -901,7 +897,6 @@ def near_text(
901897
distance: Optional[float] = None,
902898
move_to: Optional[Move] = None,
903899
move_away: Optional[Move] = None,
904-
diversity_selection: Optional[MMR] = None,
905900
) -> _HybridNearText:
906901
"""Define a near text search to be used within a hybrid query.
907902
@@ -911,7 +906,6 @@ def near_text(
911906
distance: The maximum distance to search. If not specified, the default distance specified by the server is used.
912907
move_to: Define the concepts that should be moved towards in the vector space during the search.
913908
move_away: Define the concepts that should be moved away from in the vector space during the search.
914-
diversity_selection: Apply diversity selection (e.g. MMR) to the hybrid results. Requires Weaviate >= 1.39.0.
915909
916910
Returns:
917911
A `_HybridNearText` object to be used in the `vector` parameter of the `query.hybrid` and `generate.hybrid` search methods.
@@ -922,7 +916,6 @@ def near_text(
922916
certainty=certainty,
923917
move_to=move_to,
924918
move_away=move_away,
925-
diversity_selection=diversity_selection,
926919
)
927920

928921
@staticmethod
@@ -931,14 +924,12 @@ def near_vector(
931924
*,
932925
certainty: Optional[float] = None,
933926
distance: Optional[float] = None,
934-
diversity_selection: Optional[MMR] = None,
935927
) -> _HybridNearVector:
936928
"""Define a near vector search to be used within a hybrid query.
937929
938930
Args:
939931
certainty: The minimum similarity score to return. If not specified, the default certainty specified by the server is used.
940932
distance: The maximum distance to search. If not specified, the default distance specified by the server is used.
941-
diversity_selection: Apply diversity selection (e.g. MMR) to the hybrid results. Requires Weaviate >= 1.39.0.
942933
943934
Returns:
944935
A `_HybridNearVector` object to be used in the `vector` parameter of the `query.hybrid` and `generate.hybrid` search methods.
@@ -947,7 +938,6 @@ def near_vector(
947938
vector=vector,
948939
distance=distance,
949940
certainty=certainty,
950-
diversity_selection=diversity_selection,
951941
)
952942

953943

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 & 10 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(
@@ -639,15 +640,6 @@ def _parse_hybrid(
639640

640641
near_text, near_vector, vector_bytes, vectors = None, None, None, None
641642

642-
# Hybrid diversity selection is a post-fusion, hybrid-level operation, so
643-
# it is carried on the top-level Hybrid.selection field rather than on the
644-
# near_text / near_vector sub-query.
645-
hybrid_selection = (
646-
vector.diversity_selection
647-
if isinstance(vector, (_HybridNearText, _HybridNearVector))
648-
else None
649-
)
650-
651643
if vector is None:
652644
pass
653645
elif isinstance(vector, list) and len(vector) > 0 and isinstance(vector[0], float):
@@ -748,7 +740,7 @@ def _parse_hybrid(
748740
vector_bytes=vector_bytes,
749741
vector_distance=distance,
750742
vectors=vectors,
751-
selection=self._diversity_selection_to_grpc(hybrid_selection),
743+
selection=self._diversity_selection_to_grpc(diversity_selection),
752744
bm25_search_operator=base_search_pb2.SearchOperatorOptions(
753745
operator=bm25_operator.operator,
754746
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)