diff --git a/tests/test_batchtransform.py b/tests/test_batchtransform.py index 17a9b7e3..b2a2c9c8 100644 --- a/tests/test_batchtransform.py +++ b/tests/test_batchtransform.py @@ -43,7 +43,7 @@ def return_price(data): class BatchTransformAlgorithmSetSid(TradingAlgorithm): - def initialize(self, sids): + def initialize(self, sids=None): self.history = [] self.batch_transform = return_price( @@ -111,15 +111,14 @@ class TestChangeOfSids(TestCase): ) def test_all_sids_passed(self): - algo = BatchTransformAlgorithmSetSid(self.sids, - sim_params=self.sim_params) + algo = BatchTransformAlgorithmSetSid(sim_params=self.sim_params) source = DifferentSidSource() algo.run(source) - for df, date in zip(algo.history, source.trading_days): + for i, (df, date) in enumerate(zip(algo.history, source.trading_days)): self.assertEqual(df.index[-1], date, "Newest event doesn't \ match.") - for sid in self.sids: + for sid in self.sids[:i]: self.assertIn(sid, df.columns) last_elem = len(df) - 1 diff --git a/zipline/utils/data.py b/zipline/utils/data.py index 8eead051..b12fd367 100644 --- a/zipline/utils/data.py +++ b/zipline/utils/data.py @@ -32,29 +32,61 @@ class RollingPanel(object): Restrictions: major_axis can only be a DatetimeIndex for now """ - def __init__(self, window, items, minor_axis, cap_multiple=2, + def __init__(self, window, items, sids, cap_multiple=2, dtype=np.float64): self.pos = 0 self.window = window self.items = _ensure_index(items) - self.minor_axis = _ensure_index(minor_axis) + self.minor_axis = _ensure_index(sids) self.cap_multiple = cap_multiple self.cap = cap_multiple * window self.dtype = dtype self.index_buf = np.empty(self.cap, dtype='M8[ns]') - self.buffer = pd.Panel(items=items, minor_axis=minor_axis, - major_axis=range(self.cap), - dtype=dtype) + + self.buffer = self._create_buffer() + + def _create_buffer(self): + return pd.Panel(items=self.items, minor_axis=self.minor_axis, + major_axis=range(self.cap), + dtype=self.dtype) + + def _update_buffer(self, frame): + # Drop outdated, nan-filled minors (sids) and items (fields) + non_nan_cols = set(self.buffer.dropna(axis=1).minor_axis) + new_cols = set(frame.columns) + self.minor_axis = _ensure_index(new_cols.union(non_nan_cols)) + + non_nan_items = set(self.buffer.dropna(axis=1).items) + new_items = set(frame.index) + self.items = _ensure_index(new_items.union(non_nan_items)) + + new_buffer = self._create_buffer() + # Copy old values we want to keep + # .update() is pretty slow. Ideally we would be using + # new_buffer.loc[non_nan_items, :, non_nan_cols] = + # but this triggers a bug in Pandas 0.11. Update + # this when 0.12 is released. + # https://github.com/pydata/pandas/issues/3777 + new_buffer.update( + self.buffer.loc[non_nan_items, :, non_nan_cols]) + + self.buffer = new_buffer def add_frame(self, tick, frame): """ """ if self.pos == self.cap: self._roll_data() - self.buffer.values[:, self.pos, :] = frame.ix[self.items].values + + if set(frame.columns).difference(set(self.minor_axis)) or \ + set(frame.index).difference(set(self.items)): + self._update_buffer(frame) + + self.buffer.loc[:, self.pos, :] = frame.ix[self.items].T + self.index_buf[self.pos] = tick self.pos += 1