FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

fix: Load Milvus collections once instead of on every query · feast-dev/feast@e88bb51 · GitHub

Repository navigation

Commit e88bb51

Browse files
andcommitted
fix: Load Milvus collections once instead of on every query
online_read and retrieve_online_documents_v2 called load_collection on every request, adding a round-trip to each read. The call also hid the fact that newly created collections were never loaded on Milvus servers, because indexes were created after the collection. Collections are now created with their index params, which makes Milvus load them straight away. Existing collections are loaded only if their load state isn't Loaded, and only the first time a store accesses them. The per-query load_collection calls are removed. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> Signed-off-by: Simon Hearne <simon.hearne@gmail.com>
1 parent 6b93ce9 commit e88bb51

4 files changed

Lines changed: 224 additions & 17 deletions

File tree

‎docs/reference/online-stores/milvus.md‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,13 @@ online_store:
7676

7777
The full set of configuration options is available in [MilvusOnlineStoreConfig](https://rtd.feast.dev/en/latest/#feast.infra.online_stores.milvus.MilvusOnlineStoreConfig).
7878

79+
## Collection loading
80+
81+
Feast creates collections together with their indexes, which makes Milvus load them straight away.
82+
When Feast finds an existing collection it checks its load state and loads it only if needed.
83+
Reads and searches never load collections, so a collection released outside Feast is only reloaded
84+
the next time a Feast process first accesses it.
85+
7986
## Feature views without vectors
8087

8188
Milvus requires every collection to have a vector field. For feature views that have no vector

‎sdk/python/feast/infra/online_stores/milvus_online_store/milvus.py‎

Lines changed: 20 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
FieldSchema,
1313
MilvusClient,
1414
)
15+
from pymilvus.client.types import LoadState
1516

1617
from feast import Entity
1718
from feast.feature_view import FeatureView
@@ -367,13 +368,7 @@ def _get_or_create_collection(
367368
collection_name=collection_name
368369
)
369370
if not collection_exists:
370-
self.client.create_collection(
371-
collection_name=collection_name,
372-
dimension=config.online_store.embedding_dim,
373-
schema=schema,
374-
)
375371
index_params = self.client.prepare_index_params()
376-
indices_added = False
377372
for vector_field in schema.fields:
378373
if vector_field.dtype not in [
379374
DataType.FLOAT_VECTOR,
@@ -405,19 +400,31 @@ def _get_or_create_collection(
405400
index_type="FLAT",
406401
index_name=f"vector_index_{vector_field.name}",
407402
)
408-
indices_added = True
409-
if indices_added:
410-
self.client.create_index(
411-
collection_name=collection_name,
412-
index_params=index_params,
413-
)
403+
# Every collection has at least one vector field, and every
404+
# vector field is indexed, so passing the index params here
405+
# makes Milvus create the indexes and load the collection.
406+
self.client.create_collection(
407+
collection_name=collection_name,
408+
dimension=config.online_store.embedding_dim,
409+
schema=schema,
410+
index_params=index_params,
411+
)
414412
else:
415-
self.client.load_collection(collection_name)
413+
self._ensure_loaded(collection_name)
414+
# Collections are only cached once loaded, so reads and searches
415+
# don't need to load them again.
416416
self._collections[collection_name] = self.client.describe_collection(
417417
collection_name
418418
)
419419
return self._collections[collection_name]
420420

421+
def _ensure_loaded(self, collection_name: str) -> None:
422+
"""Load an existing collection unless Milvus already has it loaded."""
423+
assert self.client is not None, "Milvus client is not initialized"
424+
load_state = self.client.get_load_state(collection_name).get("state")
425+
if load_state != LoadState.Loaded:
426+
self.client.load_collection(collection_name)
427+
421428
def online_write_batch(
422429
self,
423430
config: RepoConfig,
@@ -560,7 +567,6 @@ def online_read(
560567
+ ", ".join([f"'{e}'" for e in composite_entities])
561568
+ "]"
562569
)
563-
self.client.load_collection(collection_name)
564570
results = self.client.query(
565571
collection_name=collection_name,
566572
filter=query_filter_for_entities,
@@ -773,8 +779,6 @@ def retrieve_online_documents_v2(
773779
ann_search_field = field["name"]
774780
break
775781

776-
self.client.load_collection(collection_name)
777-
778782
if filters and filters_contain_numeric_comparison(filters):
779783
collection_field_types = {
780784
f["name"]: f["type"] for f in collection["fields"]

‎sdk/python/tests/integration/online_store/test_milvus_remote.py‎

Lines changed: 107 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from datetime import datetime, timedelta, timezone
1616
from pathlib import Path
1717
from typing import Any, Callable, Dict, Iterator, List, Optional, TypeVar
18+
from unittest.mock import patch
1819
from urllib.parse import urlparse
1920

2021
import pytest
@@ -25,7 +26,7 @@
2526
from feast.protos.feast.types.EntityKey_pb2 import EntityKey as EntityKeyProto
2627
from feast.protos.feast.types.Value_pb2 import Value as ValueProto
2728
from feast.repo_config import RepoConfig
28-
from feast.types import Float32, Int64, String
29+
from feast.types import Array, Float32, Int64, String
2930
from feast.value_type import ValueType
3031

3132
T = TypeVar("T")
@@ -171,3 +172,108 @@ def test_scalar_feature_view_round_trip(
171172
)
172173

173174
assert rows[0] is not None and rows[0]["city"].string_val == "Paris"
175+
176+
177+
def _vector_feature_view(name: str = "driver_embeddings") -> FeatureView:
178+
return FeatureView(
179+
name=name,
180+
entities=[
181+
Entity(
182+
name="driver_id", join_keys=["driver_id"], value_type=ValueType.INT64
183+
)
184+
],
185+
ttl=timedelta(days=1),
186+
schema=[
187+
Field(name="driver_id", dtype=Int64),
188+
Field(
189+
name="embedding",
190+
dtype=Array(Float32),
191+
vector_index=True,
192+
vector_search_metric="COSINE",
193+
),
194+
Field(name="city", dtype=String),
195+
],
196+
)
197+
198+
199+
def _vector_rows() -> Dict[int, Dict[str, ValueProto]]:
200+
def embedding(x: float, y: float) -> ValueProto:
201+
value = ValueProto()
202+
value.float_list_val.val.extend([x, y])
203+
return value
204+
205+
return {
206+
1: {"embedding": embedding(1.0, 0.0), "city": ValueProto(string_val="Paris")},
207+
2: {"embedding": embedding(0.0, 1.0), "city": ValueProto(string_val="Rome")},
208+
}
209+
210+
211+
def _search(
212+
store: MilvusOnlineStore,
213+
config: RepoConfig,
214+
fv: FeatureView,
215+
embedding: List[float],
216+
top_k: int = 1,
217+
**kwargs: Any,
218+
) -> List[Dict[str, ValueProto]]:
219+
results = store.retrieve_online_documents_v2(
220+
config,
221+
fv,
222+
["embedding", "city"],
223+
embedding=embedding,
224+
top_k=top_k,
225+
distance_metric="COSINE",
226+
**kwargs,
227+
)
228+
return [values for _, _, values in results if values]
229+
230+
231+
def test_load_collection_not_called_per_query(
232+
tmp_path: Path, project: str, store: MilvusOnlineStore
233+
) -> None:
234+
config = _repo_config(tmp_path, project)
235+
fv = _vector_feature_view()
236+
store.update(config, [], [fv], [], [], partial=False)
237+
_write_rows(store, config, fv, _vector_rows())
238+
239+
assert store.client is not None
240+
with patch.object(
241+
store.client, "load_collection", wraps=store.client.load_collection
242+
) as load_spy:
243+
hits = _eventually(
244+
lambda: _search(store, config, fv, [1.0, 0.0]),
245+
lambda hits: len(hits) == 1,
246+
)
247+
for _ in range(3):
248+
_read(store, config, fv, [1, 2], ["city"])
249+
_search(store, config, fv, [1.0, 0.0])
250+
251+
assert hits[0]["city"].string_val == "Paris"
252+
assert load_spy.call_count == 0
253+
254+
255+
def test_released_collection_is_loaded_once(
256+
tmp_path: Path, project: str, store: MilvusOnlineStore
257+
) -> None:
258+
config = _repo_config(tmp_path, project)
259+
fv = _vector_feature_view()
260+
store.update(config, [], [fv], [], [], partial=False)
261+
_write_rows(store, config, fv, _vector_rows())
262+
assert store.client is not None
263+
collection_name = f"{project}_{fv.name}"
264+
store.client.release_collection(collection_name)
265+
266+
# A fresh store, e.g. a new feature server process, finds it unloaded.
267+
fresh_store = MilvusOnlineStore()
268+
fresh_store.client = store.client
269+
with patch.object(
270+
store.client, "load_collection", wraps=store.client.load_collection
271+
) as load_spy:
272+
rows = _eventually(
273+
lambda: _read(fresh_store, config, fv, [1], ["city"]),
274+
lambda rows: rows[0] is not None,
275+
)
276+
_read(fresh_store, config, fv, [1], ["city"])
277+
278+
assert rows[0] is not None and rows[0]["city"].string_val == "Paris"
279+
assert load_spy.call_count == 1

‎sdk/python/tests/unit/infra/online_store/test_milvus_online_store.py‎

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,9 @@
55
from typing import Any, Dict, List, Optional
66
from unittest.mock import MagicMock, patch
77

8+
import pytest
89
from pymilvus import DataType, MilvusClient
10+
from pymilvus.client.types import LoadState
911

1012
from feast import Entity, FeatureView
1113
from feast.field import Field
@@ -211,3 +213,91 @@ def test_placeholder_values_are_finite_and_match_collection_dim(
211213

212214
data = mock_client.upsert.call_args.kwargs["data"]
213215
assert data[0][PLACEHOLDER_VECTOR_FIELD] == [0.0]
216+
217+
218+
def _existing_collection_description() -> Dict[str, Any]:
219+
return {
220+
"collection_name": "test_milvus_driver_stats",
221+
"fields": [
222+
{"name": "driver_id_pk", "type": DataType.VARCHAR, "params": {}},
223+
{"name": "driver_id", "type": DataType.VARCHAR, "params": {}},
224+
{"name": "event_ts", "type": DataType.INT64, "params": {}},
225+
{"name": "created_ts", "type": DataType.INT64, "params": {}},
226+
{"name": "trips_today", "type": DataType.VARCHAR, "params": {}},
227+
{"name": "city", "type": DataType.VARCHAR, "params": {}},
228+
{
229+
"name": PLACEHOLDER_VECTOR_FIELD,
230+
"type": DataType.FLOAT_VECTOR,
231+
"params": {"dim": PLACEHOLDER_VECTOR_DIM},
232+
},
233+
],
234+
}
235+
236+
237+
@pytest.mark.parametrize(
238+
"load_state, expected_loads",
239+
[(LoadState.Loaded, 0), (LoadState.NotLoad, 1)],
240+
)
241+
@patch(f"{MILVUS_MODULE}.MilvusClient")
242+
def test_existing_collection_is_loaded_at_most_once(
243+
mock_client_cls: MagicMock, load_state: LoadState, expected_loads: int
244+
) -> None:
245+
mock_client = _mock_client(mock_client_cls, has_collection=True)
246+
mock_client.describe_collection.return_value = _existing_collection_description()
247+
mock_client.get_load_state.return_value = {"state": load_state}
248+
mock_client.query.return_value = []
249+
250+
store = MilvusOnlineStore()
251+
config = _mock_config()
252+
fv = _scalar_feature_view()
253+
for _ in range(5):
254+
store.online_read(config, fv, [_entity_key(1)], ["city"])
255+
256+
assert mock_client.load_collection.call_count == expected_loads
257+
assert mock_client.query.call_count == 5
258+
259+
260+
@patch(f"{MILVUS_MODULE}.MilvusClient")
261+
def test_new_collection_is_created_with_indexes_so_it_loads(
262+
mock_client_cls: MagicMock,
263+
) -> None:
264+
mock_client = _mock_client(mock_client_cls, has_collection=False)
265+
266+
store = MilvusOnlineStore()
267+
store._get_or_create_collection(_mock_config(), _scalar_feature_view())
268+
269+
# MilvusClient.create_collection loads the collection when it is given
270+
# index params; creating indexes separately would leave it unloaded.
271+
assert mock_client.create_collection.call_args.kwargs["index_params"]
272+
mock_client.create_index.assert_not_called()
273+
mock_client.load_collection.assert_not_called()
274+
275+
276+
def test_load_collection_not_called_per_query(tmp_path: Path) -> None:
277+
config = _lite_config(tmp_path)
278+
fv = _scalar_feature_view()
279+
store = MilvusOnlineStore()
280+
store.update(config, [], [fv], [], [], partial=False)
281+
_write_rows(
282+
store,
283+
config,
284+
fv,
285+
{
286+
1: {
287+
"trips_today": ValueProto(float_val=1.0),
288+
"city": ValueProto(string_val="Oslo"),
289+
}
290+
},
291+
)
292+
293+
assert store.client is not None
294+
with patch.object(
295+
store.client, "load_collection", wraps=store.client.load_collection
296+
) as load_spy:
297+
for _ in range(3):
298+
_read(store, config, fv, [1], ["city"])
299+
store.retrieve_online_documents_v2(
300+
config, fv, ["city"], embedding=None, top_k=1, query_string="Oslo"
301+
)
302+
303+
assert load_spy.call_count == 0

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL