-
Notifications
You must be signed in to change notification settings - Fork 1.5k
feat: Offline Store historical features retrieval based on datetime range in Ray #5738
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 1 commit
1b5aebe
de8b2c5
c25280c
0c20ac0
f589956
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
…ange in Ray Signed-off-by: Aniket Paluskar <apaluska@redhat.com>
- Loading branch information
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1,7 +1,7 @@ | ||
| import logging | ||
| import os | ||
| import uuid | ||
| from datetime import datetime | ||
| from datetime import datetime, timedelta | ||
| from pathlib import Path | ||
| from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union | ||
|
|
||
|
|
@@ -1197,6 +1197,146 @@ def schema(self) -> pa.Schema: | |
| return pa.Table.from_pandas(df).schema | ||
|
|
||
|
|
||
| def _compute_non_entity_dates_ray( | ||
| feature_views: List[FeatureView], | ||
| start_date_opt: Optional[datetime], | ||
| end_date_opt: Optional[datetime], | ||
| ) -> Tuple[datetime, datetime]: | ||
| # Why: derive bounded time window when no entity_df is provided using explicit dates or max TTL fallback | ||
| end_date = make_tzaware(end_date_opt) if end_date_opt else make_tzaware(datetime.utcnow()) | ||
| if start_date_opt is None: | ||
| max_ttl_seconds = 0 | ||
| for fv in feature_views: | ||
| if getattr(fv, "ttl", None): | ||
| try: | ||
| ttl_val = fv.ttl | ||
| if isinstance(ttl_val, timedelta): | ||
| max_ttl_seconds = max(max_ttl_seconds, int(ttl_val.total_seconds())) | ||
| except Exception: | ||
| pass | ||
| start_date = ( | ||
| end_date - timedelta(seconds=max_ttl_seconds) | ||
| if max_ttl_seconds > 0 | ||
| else end_date - timedelta(days=30) | ||
| ) | ||
| else: | ||
| start_date = make_tzaware(start_date_opt) | ||
| return start_date, end_date | ||
|
|
||
|
|
||
| def _make_filter_range(timestamp_field: str, start_date: datetime, end_date: datetime): | ||
| # Why: factory function for time-range filtering in Ray map_batches | ||
| def _filter_range(batch: pd.DataFrame) -> pd.Series: | ||
| ts = pd.to_datetime(batch[timestamp_field], utc=True) | ||
| return (ts >= start_date) & (ts <= end_date) | ||
|
|
||
| return _filter_range | ||
|
|
||
|
|
||
| def _make_select_distinct_keys(join_keys: List[str]): | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think we should not drop rows with duplicate IDs, because there could be multiple transactions per ID and we need to choose the row based on timestamp while joining the colums from another table/view. I think this is the same case with your spark PR. Please check the postgres implementation to understand the case. Or Am I misreading this ?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Testing the case after discussion
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Previously, when entity_df=None was passed to get_historical_features(), the Ray offline store would extract only distinct entity keys and assign a single fixed timestamp (end_date) to all entities. This broke point-in-time joins for cases where multiple transactions exist per entity ID in date-time range. Now extracts distinct (entity_keys, event_timestamp) combinations, aligning with Postgres based offline store's behaviour. |
||
| # Why: factory function for distinct key projection in Ray map_batches | ||
| def _select_distinct_keys(batch: pd.DataFrame) -> pd.DataFrame: | ||
| cols = [c for c in join_keys if c in batch.columns] | ||
| if not cols: | ||
| return pd.DataFrame(columns=join_keys) | ||
| return batch[cols].drop_duplicates().reset_index(drop=True) | ||
|
|
||
| return _select_distinct_keys | ||
|
|
||
|
|
||
| def _distinct_entities_for_feature_view_ray( | ||
| store: "RayOfflineStore", | ||
| config: RepoConfig, | ||
| fv: FeatureView, | ||
| registry: BaseRegistry, | ||
| project: str, | ||
| start_date: datetime, | ||
| end_date: datetime, | ||
| ) -> Tuple[Dataset, List[str]]: | ||
| # Why: read minimal columns, filter by time, and project distinct join keys per FeatureView | ||
| ray_wrapper = get_ray_wrapper() | ||
| entities = fv.entities or [] | ||
| entity_objs = [registry.get_entity(e, project) for e in entities] | ||
| original_join_keys, _rev_feats, timestamp_field, _created_col = _get_column_names( | ||
| fv, entity_objs | ||
| ) | ||
|
|
||
| source_info = resolve_feature_view_source_with_fallback( | ||
| fv, config, is_materialization=False | ||
| ) | ||
| source_path = store._get_source_path(source_info.data_source, config) | ||
| required_columns = list(set(original_join_keys + [timestamp_field])) | ||
| ds = ray_wrapper.read_parquet(source_path, columns=required_columns) | ||
|
|
||
| field_mapping = getattr(fv.batch_source, "field_mapping", None) | ||
| if field_mapping: | ||
| ds = apply_field_mapping(ds, field_mapping) | ||
| original_join_keys = [field_mapping.get(k, k) for k in original_join_keys] | ||
| timestamp_field = field_mapping.get(timestamp_field, timestamp_field) | ||
|
|
||
| if fv.projection.join_key_map: | ||
| join_keys = [ | ||
| fv.projection.join_key_map.get(key, key) for key in original_join_keys | ||
| ] | ||
| else: | ||
| join_keys = original_join_keys | ||
|
|
||
| ds = ensure_timestamp_compatibility(ds, [timestamp_field]) | ||
| ds = ds.filter(_make_filter_range(timestamp_field, start_date, end_date)) | ||
| ds = ds.map_batches(_make_select_distinct_keys(join_keys), batch_format="pandas") | ||
| return ds, join_keys | ||
|
|
||
|
|
||
| def _make_align_columns(all_join_keys: List[str]): | ||
| # Why: factory function for schema alignment in Ray map_batches | ||
| def _align_columns(batch: pd.DataFrame) -> pd.DataFrame: | ||
| batch = batch.copy() | ||
| for k in all_join_keys: | ||
| if k not in batch.columns: | ||
| batch[k] = pd.NA | ||
| return batch[all_join_keys] | ||
|
|
||
| return _align_columns | ||
|
|
||
|
|
||
| def _make_distinct_by_keys(keys: List[str]): | ||
| # Why: factory function for deduplication in Ray map_batches | ||
| def _distinct(batch: pd.DataFrame) -> pd.DataFrame: | ||
| return batch.drop_duplicates(subset=keys).reset_index(drop=True) | ||
|
|
||
| return _distinct | ||
|
|
||
|
|
||
| def _align_and_union_entities_ray( | ||
| datasets: List[Dataset], | ||
| all_join_keys: List[str], | ||
| ) -> Dataset: | ||
| # Why: align schemas across FeatureViews and union to a unified entity set | ||
| ray_wrapper = get_ray_wrapper() | ||
| if not datasets: | ||
| return ray_wrapper.from_pandas(pd.DataFrame(columns=all_join_keys)) | ||
|
|
||
| aligned = [ | ||
| ds.map_batches(_make_align_columns(all_join_keys), batch_format="pandas") | ||
| for ds in datasets | ||
| ] | ||
| entity_ds = aligned[0] | ||
| for ds in aligned[1:]: | ||
| entity_ds = entity_ds.union(ds) | ||
| return entity_ds.map_batches(_make_distinct_by_keys(all_join_keys), batch_format="pandas") | ||
|
|
||
|
|
||
| def _add_asof_ts_ray(ds: Dataset, end_date: datetime) -> Dataset: | ||
| # Why: use a stable as-of timestamp for PIT joins when deriving entities | ||
| def _add_asof_ts(batch: pd.DataFrame) -> pd.DataFrame: | ||
| batch = batch.copy() | ||
| batch["event_timestamp"] = end_date | ||
| return batch | ||
|
|
||
| ds = ds.map_batches(_add_asof_ts, batch_format="pandas") | ||
| return ensure_timestamp_compatibility(ds, ["event_timestamp"]) | ||
|
|
||
|
|
||
| class RayOfflineStore(OfflineStore): | ||
| def __init__(self) -> None: | ||
| self._staging_location: Optional[str] = None | ||
|
|
@@ -1874,17 +2014,36 @@ def get_historical_features( | |
| config: RepoConfig, | ||
| feature_views: List[FeatureView], | ||
| feature_refs: List[str], | ||
| entity_df: Union[pd.DataFrame, str], | ||
| entity_df: Optional[Union[pd.DataFrame, str]], | ||
| registry: BaseRegistry, | ||
| project: str, | ||
| full_feature_names: bool = False, | ||
| **kwargs: Any, | ||
| ) -> RetrievalJob: | ||
| store = RayOfflineStore() | ||
| store._init_ray(config) | ||
|
|
||
| # Load entity_df as Ray dataset for distributed processing | ||
| # Load or derive entity dataset for distributed processing | ||
| ray_wrapper = get_ray_wrapper() | ||
| if isinstance(entity_df, str): | ||
| if entity_df is None: | ||
| # Non-entity mode: derive entity set from feature sources within a bounded time window | ||
| start_date, end_date = _compute_non_entity_dates_ray( | ||
| feature_views, kwargs.get("start_date"), kwargs.get("end_date") | ||
| ) | ||
| per_view_entity_ds: List[Dataset] = [] | ||
| all_join_keys: List[str] = [] | ||
| for fv in feature_views: | ||
| ds, join_keys = _distinct_entities_for_feature_view_ray( | ||
| store, config, fv, registry, project, start_date, end_date | ||
| ) | ||
| per_view_entity_ds.append(ds) | ||
| for k in join_keys: | ||
| if k not in all_join_keys: | ||
| all_join_keys.append(k) | ||
| entity_ds = _align_and_union_entities_ray(per_view_entity_ds, all_join_keys) | ||
| entity_ds = _add_asof_ts_ray(entity_ds, end_date) | ||
| entity_df_sample = entity_ds.limit(1000).to_pandas() | ||
| elif isinstance(entity_df, str): | ||
| entity_ds = ray_wrapper.read_csv(entity_df) | ||
| entity_df_sample = entity_ds.limit(1000).to_pandas() | ||
| else: | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
I think we should have make a common utility function for this, so that it can be used in all stores without repeating the code.
wdyt ?