mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-05 12:50:21 +08:00
Merge changes to RollingPanel that allow new sids to be added mid-run.
This commit is contained in:
@@ -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
|
||||
|
||||
+38
-6
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user