Skip to content

Commit 622a1c5

Browse files
committed
sdks/python: address Danny's feedback (2)
1 parent a3dedc9 commit 622a1c5

4 files changed

Lines changed: 94 additions & 65 deletions

File tree

sdks/python/apache_beam/examples/snippets/transforms/elementwise/enrichment.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -170,10 +170,17 @@ def enrichment_with_milvus():
170170
output_fields=["id", "content", "domain", "cost", "metadata"],
171171
round_decimal=2)
172172

173+
# MilvusCollectionLoadParameters is optional and provides fine-grained control
174+
# over how collections are loaded into memory. For simple use cases or when
175+
# getting started, this parameter can be omitted to use default loading
176+
# behavior. Consider using it in resource-constrained environments to optimize
177+
# memory usage and query performance.
173178
collection_load_parameters = MilvusCollectionLoadParameters()
174179

175180
milvus_search_handler = MilvusSearchEnrichmentHandler(
176-
connection_parameters, search_parameters, collection_load_parameters)
181+
connection_parameters=connection_parameters,
182+
search_parameters=search_parameters,
183+
collection_load_parameters=collection_load_parameters)
177184
with beam.Pipeline() as p:
178185
_ = (
179186
p

sdks/python/apache_beam/ml/rag/enrichment/milvus_search.py

Lines changed: 57 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -229,28 +229,41 @@ def __post_init__(self):
229229

230230

231231
@dataclass
232-
class HybridSearchNamespace:
233-
"""Namespace containing all parameters for hybrid search operations.
232+
class HybridSearchParameters:
233+
"""Parameters for hybrid (vector + keyword) search operations.
234234
235235
Args:
236236
vector: Parameters for the vector search component.
237237
keyword: Parameters for the keyword search component.
238-
hybrid: Parameters for combining the vector and keyword results.
238+
ranker: Ranker for combining vector and keyword search results.
239+
Example: RRFRanker(k=100).
240+
limit: Maximum number of results to return per query. Defaults to 3 search
241+
results.
242+
kwargs: Optional keyword arguments for additional hybrid search parameters.
243+
Enables forward compatibility.
239244
"""
240245
vector: VectorSearchParameters
241246
keyword: KeywordSearchParameters
242-
hybrid: HybridSearchParameters
247+
ranker: MilvusBaseRanker
248+
limit: int = 3
249+
kwargs: Dict[str, Any] = field(default_factory=dict)
243250

244251
def __post_init__(self):
245-
if not self.vector or not self.keyword or not self.hybrid:
252+
if not self.vector or not self.keyword:
246253
raise ValueError(
247-
"Vector, keyword, and hybrid search parameters must be provided for "
254+
"Vector and keyword search parameters must be provided for "
248255
"hybrid search")
249256

257+
if not self.ranker:
258+
raise ValueError("Ranker must be provided for hybrid search")
259+
260+
if self.limit <= 0:
261+
raise ValueError(f"Search limit must be positive, got {self.limit}")
262+
250263

251264
SearchStrategyType = Union[VectorSearchParameters,
252265
KeywordSearchParameters,
253-
HybridSearchNamespace]
266+
HybridSearchParameters]
254267

255268

256269
@dataclass
@@ -315,6 +328,22 @@ class MilvusCollectionLoadParameters:
315328
kwargs: Dict[str, Any] = field(default_factory=dict)
316329

317330

331+
@dataclass
332+
class MilvusSearchResult:
333+
"""Search result from Milvus per chunk.
334+
335+
Args:
336+
id: List of entity IDs returned from the search. Can be either string or
337+
integer IDs.
338+
distance: List of distances/similarity scores for each returned entity.
339+
fields: List of dictionaries containing additional field values for each
340+
entity. Each dictionary corresponds to one returned entity.
341+
"""
342+
id: List[Union[str, int]] = field(default_factory=list)
343+
distance: List[float] = field(default_factory=list)
344+
fields: List[Dict[str, Any]] = field(default_factory=list)
345+
346+
318347
InputT, OutputT = Union[Chunk, List[Chunk]], List[Tuple[Chunk, Dict[str, Any]]]
319348

320349

@@ -343,8 +372,8 @@ def __init__(
343372
self,
344373
connection_parameters: MilvusConnectionParameters,
345374
search_parameters: MilvusSearchParameters,
346-
collection_load_parameters: MilvusCollectionLoadParameters,
347375
*,
376+
collection_load_parameters: Optional[MilvusCollectionLoadParameters],
348377
min_batch_size: int = 1,
349378
max_batch_size: int = 1000,
350379
**kwargs):
@@ -360,7 +389,7 @@ def __init__(
360389
milvus_handler = MilvusSearchEnrichmentHandler(
361390
connection_paramters,
362391
search_parameters,
363-
collection_load_parameters,
392+
collection_load_parameters=collection_load_parameters,
364393
min_batch_size=10,
365394
max_batch_size=100)
366395
@@ -371,8 +400,8 @@ def __init__(
371400
search_parameters (MilvusSearchParameters): Configuration for search
372401
operations, including collection name, search strategy, and output
373402
fields.
374-
collection_load_parameters (MilvusCollectionLoadParameters): Parameters
375-
controlling how collections are loaded into memory, which can
403+
collection_load_parameters (Optional[MilvusCollectionLoadParameters]):
404+
Parameters controlling how collections are loaded into memory, which can
376405
significantly impact resource usage and performance.
377406
min_batch_size (int): Minimum number of elements to batch together when
378407
querying Milvus. Default is 1 (no batching when max_batch_size is 1).
@@ -390,22 +419,25 @@ def __init__(
390419
self._connection_parameters = connection_parameters
391420
self._search_parameters = search_parameters
392421
self._collection_load_parameters = collection_load_parameters
393-
self.kwargs = kwargs
422+
if not self._collection_load_parameters:
423+
self._collection_load_parameters = MilvusCollectionLoadParameters()
394424
self._batching_kwargs = {
395425
'min_batch_size': min_batch_size, 'max_batch_size': max_batch_size
396426
}
427+
self.kwargs = kwargs
397428
self.join_fn = join_fn
398429
self.use_custom_types = True
399430

400431
def __enter__(self):
401-
connectionParams = unpack_dataclass_with_kwargs(self._connection_parameters)
402-
loadCollectionParams = unpack_dataclass_with_kwargs(
432+
connection_params = unpack_dataclass_with_kwargs(
433+
self._connection_parameters)
434+
collection_load_params = unpack_dataclass_with_kwargs(
403435
self._collection_load_parameters)
404-
self._client = MilvusClient(**connectionParams)
436+
self._client = MilvusClient(**connection_params)
405437
self._client.load_collection(
406438
collection_name=self.collection_name,
407439
partition_names=self.partition_names,
408-
**loadCollectionParams)
440+
**collection_load_params)
409441

410442
def __call__(self, request: Union[Chunk, List[Chunk]], *args,
411443
**kwargs) -> List[Tuple[Chunk, Dict[str, Any]]]:
@@ -414,18 +446,18 @@ def __call__(self, request: Union[Chunk, List[Chunk]], *args,
414446
return self._get_call_response(reqs, search_result)
415447

416448
def _search_documents(self, chunks: List[Chunk]):
417-
if isinstance(self.search_strategy, HybridSearchNamespace):
449+
if isinstance(self.search_strategy, HybridSearchParameters):
418450
data = self._get_hybrid_search_data(chunks)
419-
hybrid_search_params = unpack_dataclass_with_kwargs(
420-
self.search_strategy.hybrid)
421451
return self._client.hybrid_search(
422452
collection_name=self.collection_name,
423453
partition_names=self.partition_names,
424454
output_fields=self.output_fields,
425455
timeout=self.timeout,
426456
round_decimal=self.round_decimal,
427457
reqs=data,
428-
**hybrid_search_params)
458+
ranker=self.search_strategy.ranker,
459+
limit=self.search_strategy.limit,
460+
**self.search_strategy.kwargs)
429461
elif isinstance(self.search_strategy, VectorSearchParameters):
430462
data = list(map(self._get_vector_search_data, chunks))
431463
vector_search_params = unpack_dataclass_with_kwargs(self.search_strategy)
@@ -497,14 +529,14 @@ def _get_call_response(
497529
for i in range(len(chunks)):
498530
chunk = chunks[i]
499531
hits: Hits = search_result[i]
500-
result = defaultdict(list)
532+
result = MilvusSearchResult()
501533
for i in range(len(hits)):
502534
hit: Hit = hits[i]
503535
normalized_fields = self._normalize_milvus_fields(hit.fields)
504-
result["id"].append(hit.id)
505-
result["distance"].append(hit.distance)
506-
result["fields"].append(normalized_fields)
507-
response.append((chunk, result))
536+
result.id.append(hit.id)
537+
result.distance.append(hit.distance)
538+
result.fields.append(normalized_fields)
539+
response.append((chunk, result.__dict__))
508540
return response
509541

510542
def _normalize_milvus_fields(self, fields: Dict[str, Any]):

sdks/python/apache_beam/ml/rag/enrichment/milvus_search_it_test.py

Lines changed: 12 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,6 @@
6363
MilvusCollectionLoadParameters,
6464
VectorSearchParameters,
6565
KeywordSearchParameters,
66-
HybridSearchNamespace,
6766
HybridSearchParameters,
6867
VectorSearchMetrics,
6968
KeywordSearchMetrics)
@@ -523,7 +522,9 @@ def test_invalid_query_on_non_existent_collection(self):
523522
collection_load_parameters = MilvusCollectionLoadParameters()
524523

525524
handler = MilvusSearchEnrichmentHandler(
526-
self._connection_params, search_parameters, collection_load_parameters)
525+
self._connection_params,
526+
search_parameters,
527+
collection_load_parameters=collection_load_parameters)
527528

528529
with self.assertRaises(Exception) as context:
529530
with TestPipeline() as p:
@@ -549,7 +550,9 @@ def test_invalid_query_on_non_existent_field(self):
549550
collection_load_parameters = MilvusCollectionLoadParameters()
550551

551552
handler = MilvusSearchEnrichmentHandler(
552-
self._connection_params, search_parameters, collection_load_parameters)
553+
self._connection_params,
554+
search_parameters,
555+
collection_load_parameters=collection_load_parameters)
553556

554557
with self.assertRaises(Exception) as context:
555558
with TestPipeline() as p:
@@ -569,7 +572,9 @@ def test_empty_input_chunks(self):
569572
collection_load_parameters = MilvusCollectionLoadParameters()
570573

571574
handler = MilvusSearchEnrichmentHandler(
572-
self._connection_params, search_parameters, collection_load_parameters)
575+
self._connection_params,
576+
search_parameters,
577+
collection_load_parameters=collection_load_parameters)
573578

574579
expected_chunks = []
575580

@@ -1187,16 +1192,14 @@ def test_hybrid_search(self):
11871192
search_params=addition_keyword_search_params)
11881193

11891194
hybrid_search_parameters = HybridSearchParameters(
1190-
ranker=RRFRanker(1), limit=1)
1191-
1192-
hybrid_search_ns = HybridSearchNamespace(
11931195
vector=vector_search_parameters,
11941196
keyword=keyword_search_parameters,
1195-
hybrid=hybrid_search_parameters)
1197+
ranker=RRFRanker(1),
1198+
limit=1)
11961199

11971200
search_parameters = MilvusSearchParameters(
11981201
collection_name=MILVUS_IT_CONFIG["collection_name"],
1199-
search_strategy=hybrid_search_ns,
1202+
search_strategy=hybrid_search_parameters,
12001203
output_fields=["id", "content", "metadata"],
12011204
round_decimal=1)
12021205

@@ -1306,10 +1309,6 @@ def assert_chunks_equivalent(
13061309
assert 'enrichment_data' in actual.metadata, err_msg
13071310

13081311
# For enrichment_data, ensure consistent ordering of results.
1309-
# If "expected" has values for enrichment_data but actual doesn't, that's
1310-
# acceptable since vector search results can vary based on many factors
1311-
# including implementation details, vector database state, and slight
1312-
# variations in similarity calculations.
13131312
actual_data = actual.metadata['enrichment_data']
13141313
expected_data = expected.metadata['enrichment_data']
13151314

sdks/python/apache_beam/ml/rag/enrichment/milvus_search_test.py

Lines changed: 17 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,6 @@
3030
VectorSearchParameters,
3131
KeywordSearchParameters,
3232
HybridSearchParameters,
33-
HybridSearchNamespace,
3433
MilvusBaseRanker,
3534
unpack_dataclass_with_kwargs)
3635
except ImportError as e:
@@ -285,47 +284,39 @@ class TestMilvusHybridSearchEnrichment(unittest.TestCase):
285284
"""Tests specific to hybrid search functionality"""
286285

287286
@parameterized.expand([
288-
# Missing vector in hybrid search namespace.
287+
# Missing vector in hybrid search parameters.
289288
(
290-
lambda: HybridSearchNamespace(
289+
lambda: HybridSearchParameters(
291290
vector=None, # type: ignore[arg-type]
292291
keyword=KeywordSearchParameters(anns_field="sparse_embedding"),
293-
hybrid=HybridSearchParameters(ranker=MockRanker())),
294-
"Vector, keyword, and hybrid search parameters must be provided for "
295-
"hybrid search"
292+
ranker=MockRanker()),
293+
"Vector and keyword search parameters must be provided for hybrid "
294+
"search"
296295
),
297-
# Missing keyword in hybrid search namespace.
296+
# Missing keyword in hybrid search parameters.
298297
(
299-
lambda: HybridSearchNamespace(
298+
lambda: HybridSearchParameters(
300299
vector=VectorSearchParameters(anns_field="embedding"),
301300
keyword=None, # type: ignore[arg-type]
302-
hybrid=HybridSearchParameters(ranker=MockRanker())),
303-
"Vector, keyword, and hybrid search parameters must be provided for "
304-
"hybrid search"
305-
),
306-
# Missing hybrid in hybrid search namespace.
307-
(
308-
lambda: HybridSearchNamespace(
309-
vector=VectorSearchParameters(anns_field="embedding"),
310-
keyword=KeywordSearchParameters(anns_field="sparse_embedding"),
311-
hybrid=None), # type: ignore[arg-type]
312-
"Vector, keyword, and hybrid search parameters must be provided for "
313-
"hybrid search"
301+
ranker=MockRanker()),
302+
"Vector and keyword search parameters must be provided for hybrid "
303+
"search"
314304
),
315305
# Missing ranker in hybrid search parameters.
316306
(
317-
lambda: HybridSearchNamespace(
307+
lambda: HybridSearchParameters(
318308
vector=VectorSearchParameters(anns_field="embedding"),
319309
keyword=KeywordSearchParameters(anns_field="sparse_embedding"),
320-
hybrid=HybridSearchParameters(ranker=None)), # type: ignore[arg-type]
310+
ranker=None), # type: ignore[arg-type]
321311
"Ranker must be provided for hybrid search"
322312
),
323313
# Negative limit in hybrid search parameters.
324314
(
325-
lambda: HybridSearchNamespace(
315+
lambda: HybridSearchParameters(
326316
vector=VectorSearchParameters(anns_field="embedding"),
327317
keyword=KeywordSearchParameters(anns_field="sparse_embedding"),
328-
hybrid=HybridSearchParameters(ranker=MockRanker(), limit=-1)),
318+
ranker=MockRanker(),
319+
limit=-1),
329320
"Search limit must be positive, got -1"
330321
),
331322
])
@@ -334,10 +325,10 @@ def test_invalid_search_parameters(self, create_params, expected_error_msg):
334325
with self.assertRaises(ValueError) as context:
335326
connection_params = MilvusConnectionParameters(
336327
uri="http://localhost:19530")
337-
hybrid_search_namespace = create_params()
328+
hybrid_search_params = create_params()
338329
search_params = MilvusSearchParameters(
339330
collection_name="test_collection",
340-
search_strategy=hybrid_search_namespace)
331+
search_strategy=hybrid_search_params)
341332
collection_load_params = MilvusCollectionLoadParameters()
342333

343334
_ = MilvusSearchEnrichmentHandler(

0 commit comments

Comments
 (0)