| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 6b93ce9 commit e88bb51
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -76,6 +76,13 @@ online_store: | |||
| 76 | 76 | ||
| 77 | 77 | The full set of configuration options is available in [MilvusOnlineStoreConfig](https://rtd.feast.dev/en/latest/#feast.infra.online_stores.milvus.MilvusOnlineStoreConfig). | |
| 78 | 78 | ||
| 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 | + | ||
| 79 | 86 | ## Feature views without vectors | |
| 80 | 87 | ||
| 81 | 88 | Milvus requires every collection to have a vector field. For feature views that have no vector | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,6 +12,7 @@ | |||
| 12 | 12 | FieldSchema, | |
| 13 | 13 | MilvusClient, | |
| 14 | 14 | ) | |
| 15 | + from pymilvus.client.types import LoadState | ||
| 15 | 16 | ||
| 16 | 17 | from feast import Entity | |
| 17 | 18 | from feast.feature_view import FeatureView | |
@@ -367,13 +368,7 @@ def _get_or_create_collection( | |||
| 367 | 368 | collection_name=collection_name | |
| 368 | 369 | ) | |
| 369 | 370 | 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 | - ) | ||
| 375 | 371 | index_params = self.client.prepare_index_params() | |
| 376 | - indices_added = False | ||
| 377 | 372 | for vector_field in schema.fields: | |
| 378 | 373 | if vector_field.dtype not in [ | |
| 379 | 374 | DataType.FLOAT_VECTOR, | |
@@ -405,19 +400,31 @@ def _get_or_create_collection( | |||
| 405 | 400 | index_type="FLAT", | |
| 406 | 401 | index_name=f"vector_index_{vector_field.name}", | |
| 407 | 402 | ) | |
| 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 | + ) | ||
| 414 | 412 | 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. | ||
| 416 | 416 | self._collections[collection_name] = self.client.describe_collection( | |
| 417 | 417 | collection_name | |
| 418 | 418 | ) | |
| 419 | 419 | return self._collections[collection_name] | |
| 420 | 420 | ||
| 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 | + | ||
| 421 | 428 | def online_write_batch( | |
| 422 | 429 | self, | |
| 423 | 430 | config: RepoConfig, | |
@@ -560,7 +567,6 @@ def online_read( | |||
| 560 | 567 | + ", ".join([f"'{e}'" for e in composite_entities]) | |
| 561 | 568 | + "]" | |
| 562 | 569 | ) | |
| 563 | - self.client.load_collection(collection_name) | ||
| 564 | 570 | results = self.client.query( | |
| 565 | 571 | collection_name=collection_name, | |
| 566 | 572 | filter=query_filter_for_entities, | |
@@ -773,8 +779,6 @@ def retrieve_online_documents_v2( | |||
| 773 | 779 | ann_search_field = field["name"] | |
| 774 | 780 | break | |
| 775 | 781 | ||
| 776 | - self.client.load_collection(collection_name) | ||
| 777 | - | ||
| 778 | 782 | if filters and filters_contain_numeric_comparison(filters): | |
| 779 | 783 | collection_field_types = { | |
| 780 | 784 | f["name"]: f["type"] for f in collection["fields"] | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,6 +15,7 @@ | |||
| 15 | 15 | from datetime import datetime, timedelta, timezone | |
| 16 | 16 | from pathlib import Path | |
| 17 | 17 | from typing import Any, Callable, Dict, Iterator, List, Optional, TypeVar | |
| 18 | + from unittest.mock import patch | ||
| 18 | 19 | from urllib.parse import urlparse | |
| 19 | 20 | ||
| 20 | 21 | import pytest | |
@@ -25,7 +26,7 @@ | |||
| 25 | 26 | from feast.protos.feast.types.EntityKey_pb2 import EntityKey as EntityKeyProto | |
| 26 | 27 | from feast.protos.feast.types.Value_pb2 import Value as ValueProto | |
| 27 | 28 | from feast.repo_config import RepoConfig | |
| 28 | - from feast.types import Float32, Int64, String | ||
| 29 | + from feast.types import Array, Float32, Int64, String | ||
| 29 | 30 | from feast.value_type import ValueType | |
| 30 | 31 | ||
| 31 | 32 | T = TypeVar("T") | |
@@ -171,3 +172,108 @@ def test_scalar_feature_view_round_trip( | |||
| 171 | 172 | ) | |
| 172 | 173 | ||
| 173 | 174 | 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 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,7 +5,9 @@ | |||
| 5 | 5 | from typing import Any, Dict, List, Optional | |
| 6 | 6 | from unittest.mock import MagicMock, patch | |
| 7 | 7 | ||
| 8 | + import pytest | ||
| 8 | 9 | from pymilvus import DataType, MilvusClient | |
| 10 | + from pymilvus.client.types import LoadState | ||
| 9 | 11 | ||
| 10 | 12 | from feast import Entity, FeatureView | |
| 11 | 13 | from feast.field import Field | |
@@ -211,3 +213,91 @@ def test_placeholder_values_are_finite_and_match_collection_dim( | |||
| 211 | 213 | ||
| 212 | 214 | data = mock_client.upsert.call_args.kwargs["data"] | |
| 213 | 215 | 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 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments