mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-07 11:20:19 +08:00
ENH: Make RollingPanel update itself if new fields arrive.
Before we preinitialized the BT's fields and sids. Thus, no new ones could be added after initialization. This should be fixed now.
This commit is contained in:
committed by
Eddie Hebert
parent
33c23af503
commit
236fe92a53
@@ -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
|
||||
|
||||
+30
-5
@@ -32,28 +32,53 @@ 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
|
||||
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()
|
||||
|
||||
if set(frame.columns).difference(set(self.minor_axis)) or \
|
||||
set(frame.index).difference(set(self.items)):
|
||||
self._update_buffer(frame)
|
||||
|
||||
self.buffer.values[:, self.pos, :] = frame.ix[self.items].values
|
||||
self.index_buf[self.pos] = tick
|
||||
|
||||
|
||||
Reference in New Issue
Block a user