diff --git a/zipline/test_algorithms.py b/zipline/test_algorithms.py index 8fa80b22..e771ca5e 100644 --- a/zipline/test_algorithms.py +++ b/zipline/test_algorithms.py @@ -271,39 +271,39 @@ class BatchTransformAlgorithm(TradingAlgorithm): self.history_return_arbitrary_fields = [] self.return_price_class = ReturnPriceBatchTransform( - market_aware=False, refresh_period=self.refresh_period, - delta=timedelta(days=self.window_length) + window_length=self.window_length, + fillna=False ) self.return_price_decorator = return_price_batch_decorator( - market_aware=False, refresh_period=self.refresh_period, - delta=timedelta(days=self.window_length) + window_length=self.window_length, + fillna=False ) self.return_args_batch = return_args_batch_decorator( - market_aware=False, refresh_period=self.refresh_period, - delta=timedelta(days=self.window_length) + window_length=self.window_length, + fillna=False ) self.return_price_market_aware = ReturnPriceBatchTransform( - market_aware=True, refresh_period=self.refresh_period, - window_length=self.window_length + window_length=self.window_length, + fillna=False ) self.return_price_more_days_than_refresh = ReturnPriceBatchTransform( - market_aware=True, refresh_period=1, - window_length=3 + window_length=3, + fillna=False ) self.return_arbitrary_fields = return_data( - market_aware=True, refresh_period=1, - window_length=3 + window_length=3, + fillna=False ) self.set_slippage(FixedSlippage()) diff --git a/zipline/transforms/utils.py b/zipline/transforms/utils.py index 0050241d..1785400d 100644 --- a/zipline/transforms/utils.py +++ b/zipline/transforms/utils.py @@ -321,19 +321,21 @@ class BatchTransform(EventWindow): def __init__(self, func=None, refresh_period=None, - market_aware=True, - delta=None, - window_length=None): + window_length=None, + fillna=True, + fill_method='ffill'): - super(BatchTransform, self).__init__(market_aware, - window_length=window_length, - delta=delta) + super(BatchTransform, self).__init__(True, + window_length=window_length) if func is not None: self.compute_transform_value = func else: 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 self.trading_days_since_update = 0 @@ -437,14 +439,15 @@ class BatchTransform(EventWindow): # concatenate different sids into one df df = pd.DataFrame.from_dict(values_per_sid) - # 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='ffill') - # Drop any empty rows after the fill. - # This will drop a leading row of N/A - df = df.dropna() + + 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