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:
Thomas Wiecki
2012-12-06 14:43:20 -05:00
parent 5f6839beea
commit ba46c6292f
3 changed files with 58 additions and 85 deletions
+8 -39
View File
@@ -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
View File
@@ -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
View File
@@ -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