mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-07 11:20:19 +08:00
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.
This commit is contained in:
@@ -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.")
|
||||
|
||||
+25
-14
@@ -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):
|
||||
"""
|
||||
|
||||
+25
-32
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user