diff --git a/tests/pipeline/test_events.py b/tests/pipeline/test_events.py index 0a9c8b1d..9663b94e 100644 --- a/tests/pipeline/test_events.py +++ b/tests/pipeline/test_events.py @@ -6,6 +6,7 @@ from unittest import TestCase import blaze as bz from nose_parameterized import parameterized +import numpy as np from numpy.testing import assert_array_equal import pandas as pd from pandas.util.testing import assert_series_equal @@ -37,6 +38,7 @@ ABSTRACT_EXPECTED_COLS_ERROR = 'abstract methods event_date_col, expected_cols' class EventDataSet(DataSet): previous_announcement = Column(datetime64ns_dtype) + next_announcement = Column(datetime64ns_dtype) class EventDataSetLoader(EventsLoader): @@ -62,6 +64,12 @@ class EventDataSetLoader(EventsLoader): self.dataset.previous_announcement, ) + @lazyval + def next_announcement_loader(self): + return self._next_event_date_loader( + self.dataset.next_announcement, + ) + # Test case just for catching an error when multiple columns are in the wrong # data format, so no loader defined. @@ -90,6 +98,34 @@ dtx = pd.date_range('2014-01-01', '2014-01-10') class EventLoaderTestCase(TestCase): + def test_null_in_event_date_col(self): + # Tests that if there is a null date in the event date column, it is + # filtered out and does not break on loading the adjusted array. + dates_with_null = pd.Series(dtx) + dates_with_null[2] = pd.NaT + events_by_sid = {0: pd.DataFrame({ANNOUNCEMENT_FIELD_NAME: + dates_with_null, + TS_FIELD_NAME: dtx})} + loader = EventDataSetLoader( + dtx, + events_by_sid, + ) + + prev_result = loader.load_adjusted_array({ + EventDataSet.previous_announcement + }, dtx, [0], [True])[EventDataSet.previous_announcement].data[:, 0] + + next_result = loader.load_adjusted_array({ + EventDataSet.next_announcement + }, dtx, [0], [True])[EventDataSet.next_announcement].data[:, 0] + + expected_prev = dates_with_null[:] + expected_prev[2] = dtx[1] + assert_array_equal(prev_result, expected_prev) + expected_next = dates_with_null[:] + expected_next[2] = np.datetime64('NaT') + assert_array_equal(next_result, expected_next) + def assert_loader_error(self, events_by_sid, error, msg, infer_timestamps, loader): with self.assertRaisesRegexp(error, re.escape(msg)): @@ -261,25 +297,3 @@ class BlazeEventDataSetLoader(BlazeEventsLoader): ).__init__(expr, dataset=dataset, **kwargs) - - -class BlazeEventLoaderNullInDateColumnTestCase(TestCase): - def test_null_in_event_date_col(self): - # Tests that if there is a null date in the event date column, it is - # filtered out and does not break on loading the adjusted array. - dates_with_null = pd.Series(dtx) - dates_with_null[2] = pd.NaT - events_by_sid = pd.DataFrame({SID_FIELD_NAME: 0, - ANNOUNCEMENT_FIELD_NAME: dates_with_null, - TS_FIELD_NAME: dtx}) - loader = BlazeEventDataSetLoader( - bz.data(events_by_sid), - ) - - result = loader.load_adjusted_array({ - EventDataSet.previous_announcement - }, dtx, [0], [True])[EventDataSet.previous_announcement].data[:, 0] - - expected = dates_with_null.copy(True) - expected[2] = dtx[1] - assert_array_equal(result, expected) diff --git a/zipline/pipeline/loaders/blaze/events.py b/zipline/pipeline/loaders/blaze/events.py index 34d7f37d..6fae418f 100644 --- a/zipline/pipeline/loaders/blaze/events.py +++ b/zipline/pipeline/loaders/blaze/events.py @@ -77,11 +77,8 @@ class BlazeEventsLoader(PipelineLoader): ) expected_fields = self._expected_fields - expr = expr[list(expected_fields)] self._expr = bind_expression_to_resources( - expr[expr[ - self.concrete_loader.event_date_col - ].notnull()], + expr[list(expected_fields)], resources, ) self._odo_kwargs = odo_kwargs if odo_kwargs is not None else {} diff --git a/zipline/pipeline/loaders/consensus_estimates.py b/zipline/pipeline/loaders/consensus_estimates.py index f22378e0..9d8f20f8 100644 --- a/zipline/pipeline/loaders/consensus_estimates.py +++ b/zipline/pipeline/loaders/consensus_estimates.py @@ -152,6 +152,5 @@ class ConsensusEstimatesLoader(EventsLoader): def previous_actual_value_loader(self): return self._previous_event_value_loader( self.dataset.previous_actual_value, - RELEASE_DATE_FIELD_NAME, ACTUAL_VALUE_FIELD_NAME, ) diff --git a/zipline/pipeline/loaders/events.py b/zipline/pipeline/loaders/events.py index 42afac97..be1a8ba3 100644 --- a/zipline/pipeline/loaders/events.py +++ b/zipline/pipeline/loaders/events.py @@ -152,7 +152,8 @@ class EventsLoader(PipelineLoader): raise ValueError( WRONG_MANY_COL_DATA_FORMAT_ERROR.format(sid=k) ) - + self.events_by_sid = {sid: df.dropna(subset=[self.event_date_col]) for + sid, df in self.events_by_sid.items()} self.dataset = dataset def get_loader(self, column): diff --git a/zipline/testing/fixtures.py b/zipline/testing/fixtures.py index 3ae3deb8..d1547489 100644 --- a/zipline/testing/fixtures.py +++ b/zipline/testing/fixtures.py @@ -912,8 +912,8 @@ class WithPipelineEventDataLoader(with_metaclass( frame = pd.DataFrame({sid: get_values_for_date_ranges( zip_date_index_with_vals, vals[sid], - pd.DatetimeIndex(zip(*date_intervals[sid])[0]), - pd.DatetimeIndex(zip(*date_intervals[sid])[1]), + pd.DatetimeIndex(list(zip(*date_intervals[sid]))[0]), + pd.DatetimeIndex(list(zip(*date_intervals[sid]))[1]), dates ) for sid in self.get_sids()[:-1]}) frame[self.get_sids()[-1]] = zip_date_index_with_vals(