From ba46c6292f57932e21dfe836261cca1caf36a8e7 Mon Sep 17 00:00:00 2001 From: Thomas Wiecki Date: Thu, 6 Dec 2012 14:43:20 -0500 Subject: [PATCH] ENH: DataPanel is now created more flexible over all sids and all given fields. Added unittest to test for nan-filling. Added backwards filling by default. --- tests/test_transforms.py | 47 ++++++------------------------ zipline/test_algorithms.py | 39 ++++++++++++++++--------- zipline/transforms/utils.py | 57 ++++++++++++++++--------------------- 3 files changed, 58 insertions(+), 85 deletions(-) diff --git a/tests/test_transforms.py b/tests/test_transforms.py index cb620d84..9ed507f2 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -313,23 +313,16 @@ class TestBatchTransform(TestCase): def test_event_window(self): algo = BatchTransformAlgorithm() algo.run(self.source) - - self.assertEqual(algo.history_return_price_class[:2], - [None, None], + wl = algo.window_length + self.assertEqual(algo.history_return_price_class[:wl], + [None] * wl, "First two iterations should return None") - self.assertEqual(algo.history_return_price_decorator[:2], - [None, None], + self.assertEqual(algo.history_return_price_decorator[:wl], + [None] * wl, "First two iterations should return None") - self.assertEqual(algo.history_return_price_market_aware[:2], - [None, None], - "First two iterations should return None") - self.assertEqual(algo.history_return_more_days_than_refresh[:3], - [None, None, None], - "First five iterations should return None") self.assertTrue(isinstance( - algo.history_return_more_days_than_refresh[4], - pd.DataFrame), - "Sixth iteration should not be None" + algo.history_return_price_class[wl + 1], + pd.DataFrame) ) # Test whether arbitrary fields can be added to datapanel @@ -341,7 +334,7 @@ class TestBatchTransform(TestCase): self.assertTrue(all( field['arbitrary'].values.flatten() == - ['test'] * algo.window_length), + [123] * algo.window_length), 'arbitrary dataframe should contain only "test"' ) @@ -365,27 +358,3 @@ class TestBatchTransform(TestCase): algo.history_return_args, [None, None, None, expected_item, expected_item, expected_item]) - - -class TestBatchTransformMarketAware(TestCase): - def setUp(self): - setup_logger(self) - start = pd.datetime(1993, 1, 1, 0, 0, 0, 0, pytz.utc) - end = pd.datetime(1994, 1, 1, 0, 0, 0, 0, pytz.utc) - - self.data = factory.load_from_yahoo(stocks=['AAPL'], - indexes={}, - start=start, end=end) - - def test_event_window(self): - days = 50 - algo = BatchTransformAlgorithm(days=days, refresh_period=days) - algo.run(self.data) - - self.assertEqual(algo.history_return_price_market_aware[:days], - [None] * days, - "First {days} iterations should return None" - .format(days=days)) - self.assertFalse(algo.history_return_price_market_aware[days + 1] - is None, - "Window is contains too many Nones.") diff --git a/zipline/test_algorithms.py b/zipline/test_algorithms.py index d3959d12..21c4b62c 100644 --- a/zipline/test_algorithms.py +++ b/zipline/test_algorithms.py @@ -266,9 +266,8 @@ class BatchTransformAlgorithm(TradingAlgorithm): self.history_return_price_class = [] self.history_return_price_decorator = [] self.history_return_args = [] - self.history_return_price_market_aware = [] - self.history_return_more_days_than_refresh = [] self.history_return_arbitrary_fields = [] + self.history_return_nan = [] self.return_price_class = ReturnPriceBatchTransform( refresh_period=self.refresh_period, @@ -294,18 +293,20 @@ class BatchTransformAlgorithm(TradingAlgorithm): fillna=False ) - self.return_price_more_days_than_refresh = ReturnPriceBatchTransform( - refresh_period=1, - window_length=3, + self.return_arbitrary_fields = return_data( + refresh_period=self.refresh_period, + window_length=self.window_length, fillna=False ) - self.return_arbitrary_fields = return_data( - refresh_period=1, - window_length=3, - fillna=False + self.return_nan = ReturnPriceBatchTransform( + refresh_period=self.refresh_period, + window_length=self.window_length, + fillna=True ) + self.iter = 0 + self.set_slippage(FixedSlippage()) def handle_data(self, data): @@ -316,18 +317,28 @@ class BatchTransformAlgorithm(TradingAlgorithm): self.history_return_args.append( self.return_args_batch.handle_data( data, *self.args, **self.kwargs)) - self.history_return_price_market_aware.append( - self.return_price_market_aware.handle_data(data)) - self.history_return_more_days_than_refresh.append( - self.return_price_more_days_than_refresh.handle_data(data)) new_data = deepcopy(data) for sid in new_data: - new_data[sid]['arbitrary'] = 'test' + new_data[sid]['arbitrary'] = 123 self.history_return_arbitrary_fields.append( self.return_arbitrary_fields.handle_data(new_data)) + # nan every second event price + if self.iter % 2 == 0: + self.history_return_nan.append( + self.return_nan.handle_data(data)) + else: + nan_data = deepcopy(data) + import numpy as np + for sid in nan_data.iterkeys(): + nan_data[sid].price = np.nan + self.history_return_nan.append( + self.return_nan.handle_data(nan_data)) + + self.iter += 1 + class SetPortfolioAlgorithm(TradingAlgorithm): """ diff --git a/zipline/transforms/utils.py b/zipline/transforms/utils.py index 9fc3c4e2..e6f6a251 100644 --- a/zipline/transforms/utils.py +++ b/zipline/transforms/utils.py @@ -322,8 +322,7 @@ class BatchTransform(EventWindow): func=None, refresh_period=None, window_length=None, - fillna=True, - fill_method='ffill'): + fillna=True): super(BatchTransform, self).__init__(True, window_length=window_length) @@ -334,7 +333,6 @@ class BatchTransform(EventWindow): self.compute_transform_value = self.get_value self.fillna = fillna - self.fill_method = fill_method self.refresh_period = refresh_period self.window_length = window_length @@ -381,7 +379,7 @@ class BatchTransform(EventWindow): "Each sid must have the same keys." unwanted_fields = set(['portfolio', 'sid', 'dt', 'type', - 'datetime']) + 'datetime', 'source_id']) return sid_keys[0] - unwanted_fields def handle_add(self, event): @@ -421,38 +419,33 @@ class BatchTransform(EventWindow): """ # This Panel data structure ultimately gets passed to the # user-overloaded get_value() method. - # - # self.ticks contains ndicts with data, dt keys. - # event parameter is an ndict with data, dt keys. - fields = {} + sids = set.union(*[set(tick.data.keys()) for tick in self.ticks]) + dts = [tick.dt for tick in self.ticks] - for field_name in self.field_names: - # Extract all used sids - sids = set.union(*[set(tick.data.keys()) for tick in self.ticks]) + data = pd.Panel(items=self.field_names, major_axis=dts, + minor_axis=sids) - values_per_sid = {} + # Fill data panel + for tick in self.ticks: + dt = tick.dt + for sid, fields in tick.data.iteritems(): + for field_name in self.field_names: + data[field_name][sid].ix[dt] = fields[field_name] - for sid in sids: - values_per_sid[sid] = pd.Series( - {tick.data[sid].dt: tick.data[sid][field_name] - for tick in self.ticks} - ) + if self.fillna: + # Fills in gaps of missing data during transform + # of multiple stocks. E.g. we may be missing + # minute data because of illiquidity of one stock + data = data.fillna(method='ffill') + # Since we already forward filled, this can only + # fill the first value if it was missing. + # It's not wise to drop a complete column via dropna()) + # because of one missing value. + data = data.fillna(method='bfill') - # concatenate different sids into one df - df = pd.DataFrame.from_dict(values_per_sid) - - if self.fillna: - # Fills in gaps of missing data during transform - # of multiple stocks. E.g. we may be missing - # minute data because of illiquidity of one stock - df = df.fillna(method=self.fill_method) - # Drop any empty rows after the fill. - # This will drop a leading row of N/A - df = df.dropna() - - fields[field_name] = df - - data = pd.Panel.from_dict(fields, orient='items') + # Drop any empty rows after the fill. + # This will drop a leading row of N/A + data = data.dropna() return data