|
18 | 18 |
|
19 | 19 | import os |
20 | 20 | from datetime import datetime |
21 | | -import dateutil.parser |
22 | 21 |
|
23 | 22 | import grpc |
24 | 23 | import pandas as pd |
@@ -220,17 +219,27 @@ def get_serving_data(self, feature_set, entity_keys, ts_range=None): |
220 | 219 | representing the data wanted |
221 | 220 | entity_keys (:obj: `list` of :obj: `str): list of entity keys |
222 | 221 | ts_range (:obj: `list` of str, optional): size 2 list of start |
223 | | - timestamp and end timestamp, in ISO 8601 format. It will |
| 222 | + and end time, in datetime type. It will |
224 | 223 | filter out any feature value having event timestamp outside |
225 | 224 | of the ts_range. |
226 | 225 |
|
227 | 226 | Returns: |
228 | 227 | pandas.DataFrame: DataFrame of results |
229 | 228 | """ |
| 229 | + start = None |
| 230 | + end = None |
| 231 | + if ts_range is not None: |
| 232 | + if len(ts_range) != 2: |
| 233 | + raise ValueError("ts_range must have len 2") |
| 234 | + start = ts_range[0] |
| 235 | + end = ts_range[1] |
| 236 | + if type(start) is not datetime or type(end) is not datetime: |
| 237 | + raise TypeError("start and end must be datetime type") |
| 238 | + |
230 | 239 | request = self._build_serving_request(feature_set, entity_keys) |
231 | 240 | self._connect_serving() |
232 | 241 | return self._response_to_df(feature_set, self._serving_service_stub |
233 | | - .QueryFeatures(request), ts_range) |
| 242 | + .QueryFeatures(request), start, end) |
234 | 243 |
|
235 | 244 | def download_dataset(self, dataset_info, dest, staging_location, |
236 | 245 | file_type=FileType.CSV): |
@@ -299,22 +308,15 @@ def _build_serving_request(self, feature_set, entity_keys): |
299 | 308 | entityId=entity_keys, |
300 | 309 | featureId=feature_set.features) |
301 | 310 |
|
302 | | - def _response_to_df(self, feature_set, response, ts_range = None): |
303 | | - start = None |
304 | | - end = None |
305 | | - if ts_range is not None: |
306 | | - if len(ts_range) != 2: |
307 | | - raise ValueError("ts_range must have len 2") |
308 | | - start = dateutil.parser.parse(ts_range[0]) |
309 | | - end = dateutil.parser.parse(ts_range[1]) |
310 | | - |
| 311 | + def _response_to_df(self, feature_set, response, start=None, end=None): |
| 312 | + is_filter_time = start is not None and end is not None |
311 | 313 | df = pd.DataFrame(columns=[feature_set.entity] + feature_set.features) |
312 | 314 | for entity_id in response.entities: |
313 | 315 | feature_map = response.entities[entity_id].features |
314 | 316 | row = {response.entityName: entity_id} |
315 | 317 | for feature_id in feature_map: |
316 | 318 | v = feature_map[feature_id].value |
317 | | - if ts_range is not None and not _is_granularity_none( |
| 319 | + if is_filter_time and not _is_granularity_none( |
318 | 320 | feature_id): |
319 | 321 | ts = feature_map[feature_id].timestamp.ToDatetime() |
320 | 322 | if ts < start or ts > end: |
|
0 commit comments