ENH: Changed algorithm class to provide .run() method. Testing of multiple param combos of DMA.

This commit is contained in:
Thomas Wiecki
2012-08-12 20:24:34 -04:00
parent 51f01c0f1f
commit 3884b15eb6
5 changed files with 212 additions and 48 deletions
+9 -6
View File
@@ -167,6 +167,8 @@ class PerformanceTracker(object):
self.last_dict = None
self.exceeded_max_loss = False
self.compute_risk_metrics = True
self.results_socket = None
self.results_addr = None
@@ -297,12 +299,13 @@ class PerformanceTracker(object):
self.returns.append(todays_return_obj)
#calculate risk metrics for cumulative performance
self.cumulative_risk_metrics = risk.RiskMetrics(
start_date=self.period_start,
end_date=self.market_close.replace(hour=0, minute=0, second=0),
returns=self.returns,
trading_environment=self.trading_environment
)
if self.compute_risk_metrics:
self.cumulative_risk_metrics = risk.RiskMetrics(
start_date=self.period_start,
end_date=self.market_close.replace(hour=0, minute=0, second=0),
returns=self.returns,
trading_environment=self.trading_environment
)
# increment the day counter before we move markers forward.
self.day_count += 1.0
+1 -1
View File
@@ -16,7 +16,7 @@ def date_sorted_sources(*sources):
for source in sources:
assert iter(source), "Source %s not iterable" % source
assert source.__class__.__dict__.has_key('get_hash'), "No get_hash"
assert hasattr(source, 'get_hash'), "No get_hash"
# Get name hashes to pass to date_sort.
names = [source.get_hash() for source in sources]
+62
View File
@@ -4,6 +4,11 @@ and zipline development
"""
import random
import pytz
from copy import copy
import pandas as pd
from zipline import ndict
from zipline.protocol import DATASOURCE_TYPE
from itertools import chain, cycle, ifilter, izip
from datetime import datetime, timedelta
@@ -132,6 +137,63 @@ class SpecificEquityTrades(object):
return filtered
class DataFrameSource(SpecificEquityTrades):
"""
Yields all events in event_list that match the given sid_filter.
If no event_list is specified, generates an internal stream of events
to filter. Returns all events if filter is None.
Configuration options:
count : integer representing number of trades
sids : list of values representing simulated internal sids
start : start date
delta : timedelta between internal events
filter : filter to remove the sids
"""
def __init__(self, data, **kwargs):
assert isinstance(data.index, pd.tseries.index.DatetimeIndex)
self.data = data
# Unpack config dictionary with default values.
self.count = kwargs.get('count', 500)
self.sids = kwargs.get('sids', [1, 2])
self.start = kwargs.get('start', datetime(1957, 1, 1, 0, tzinfo = pytz.utc))
self.end = kwargs.get('end', datetime(2010, 1, 1, tzinfo=pytz.utc))
self.delta = kwargs.get('delta', timedelta(days = 1))
# Default to None for event_list and filter.
self.filter = kwargs.get('filter')
# Hash_value for downstream sorting.
self.arg_string = hash_args(data, **kwargs)
self.generator = self.create_fresh_generator()
def create_fresh_generator(self):
def _generator(df=self.data):
for dt, series in df.iterrows():
dt = dt.tz_localize('UTC')
if (self.start > dt) or (dt < self.end):
continue
event = {'dt': dt,
'source_id': self.get_hash(),
'type': DATASOURCE_TYPE.TRADE
}
for sid, price in series.iterkv():
event = copy(event)
event['sid'] = 0
event['price'] = price
yield ndict(event)
# Return the filtered event stream.
return _generator()
# !!!!!!! Deprecated for now !!!!!!!!!
def RandomEquityTrades(object):
+78 -32
View File
@@ -1,5 +1,16 @@
import pandas as pd
import numpy as np
from zipline.gens.mavg import MovingAverage
from datetime import datetime, timedelta
from zipline.finance.trading import SIMULATION_STYLE
from zipline.utils import factory
from zipline.gens.tradegens import SpecificEquityTrades, DataFrameSource
from zipline.protocol import DATASOURCE_TYPE
from zipline import ndict
from zipline.utils.factory import create_trading_environment
from zipline.gens.transform import StatefulTransform
from zipline.lines import SimulatedTrading
class BuySellAlgorithm(object):
"""Algorithm that buys and sells alternatingly. The amount for
@@ -53,12 +64,77 @@ class BuySellAlgorithm(object):
# Algorithm base class, user algorithms inherit from this as they
# don't want to have to copy and know about set_order and
# set_portfolio
class Algorithm(object):
class TradingAlgorithm(object):
def _setup(self, compute_risk_metrics=False):
assert hasattr(self, 'source'), 'source not set.'
assert hasattr(self, 'sids'), "sids not set."
environment = create_trading_environment()
# Create transforms by wrapping them into StatefulTransforms
transforms = []
for namestring, trans_descr in self.registered_transforms.iteritems():
sf = StatefulTransform(
trans_descr['class'],
*trans_descr['args'],
**trans_descr['kwargs']
)
sf.namestring = namestring
transforms.append(sf)
results_socket_uri = None
context = None
sim_id = None
style = SIMULATION_STYLE.FIXED_SLIPPAGE
self.simulated_trading = SimulatedTrading(
[self.source],
transforms,
self,
environment,
style,
results_socket_uri,
context,
sim_id)
#self.simulated_trading.trading_client.performance_tracker.compute_risk_metrics = compute_risk_metrics
def _create_daily_stats(self, perfs):
# create daily stats dataframe
daily_perfs = []
cum_perfs = []
for perf in perfs:
if 'daily_perf' in perf:
daily_perfs.append(perf['daily_perf'])
else:
cum_perfs.append(perf)
daily_dts = [np.datetime64(perf['period_close'], utc=True) for perf in daily_perfs]
daily_stats = pd.DataFrame(daily_perfs, index=daily_dts)
return daily_stats
def run(self, data, compute_risk_metrics=False):
self.source = DataFrameSource(data, sids=self.sids)
self._setup(compute_risk_metrics=compute_risk_metrics)
# drain simulated_trading
perfs = [perf for perf in self.simulated_trading]
daily_stats = self._create_daily_stats(perfs)
return daily_stats
def set_portfolio(self, portfolio):
self.portfolio = portfolio
def set_order(self, order_callable):
self.order = order_callable
def get_sid_filter(self):
return [self.sid]
return self.sids
def set_logger(self, logger):
self.logger = logger
@@ -73,33 +149,3 @@ class Algorithm(object):
self.registered_transforms[tag] = {'class': transform_class,
'args': args,
'kwargs': kwargs}
# Inherits from Algorithm base class
class DMA(Algorithm):
"""Dual Moving Average algorithm.
"""
def __init__(self, sid, amount, short_window=20, long_window=40):
self.sid = sid
self.amount = amount
self.done = False
self.order = None
self.frame_count = 0
self.portfolio = None
self.orders = []
self.market_entered = False
self.prices = []
self.events = 0
self.add_transform(MovingAverage, 'short_mavg', ['price'], market_aware=False, delta=timedelta(days=short_window))
self.add_transform(MovingAverage, 'long_mavg', ['price'], market_aware=False, delta=timedelta(days=long_window))
def handle_data(self, data):
self.events += 1
# access transforms via their user-defined tag
if (data[self.sid].short_mavg > data[self.sid].long_mavg) and not self.market_entered:
self.order(self.sid, 100)
self.market_entered = True
elif (data[self.sid].short_mavg < data[self.sid].long_mavg) and self.market_entered:
self.order(self.sid, -100)
self.market_entered = False
+62 -9
View File
@@ -1,19 +1,72 @@
from zipline.lines import Zipline
from zipline.optimize.algorithms import DMA
import pandas as pd
import numpy as np
#from mpl_toolkits.mplot3d import Axes3D
import matplotlib.pyplot as plt
import cProfile
from zipline.gens.mavg import MovingAverage
from zipline.optimize.algorithms import TradingAlgorithm
from datetime import timedelta
def run():
myalgo = DMA(sid=0, amount=100)
zp = Zipline(algorithm=myalgo, sources='S&P')
stats = zp.run()
print stats
from mpi4py_map import map
# Inherits from Algorithm base class
class DMA(TradingAlgorithm):
"""Dual Moving Average algorithm.
"""
def __init__(self, sid, amount=100, short_window=20, long_window=40):
self.sids = [sid]
self.amount = amount
self.done = False
self.order = None
self.frame_count = 0
self.portfolio = None
self.orders = []
self.market_entered = False
self.prices = []
self.events = 0
self.add_transform(MovingAverage, 'short_mavg', ['price'],
market_aware=False,
delta=timedelta(days=int(short_window)))
self.add_transform(MovingAverage, 'long_mavg', ['price'],
market_aware=False,
delta=timedelta(days=int(long_window)))
def handle_data(self, data):
self.events += 1
sid = self.sids[0]
# access transforms via their user-defined tag
if (data[sid].short_mavg > data[sid].long_mavg) and not self.market_entered:
self.order(sid, 100)
self.market_entered = True
elif (data[sid].short_mavg < data[sid].long_mavg) and self.market_entered:
self.order(sid, -100)
self.market_entered = False
def run((short_window, long_window)):
data = pd.DataFrame.from_csv('SP500.csv')
myalgo = DMA(sid=0, amount=100, short_window=short_window, long_window=long_window)
stats = myalgo.run(data, compute_risk_metrics=False)
stats['sw'] = short_window
stats['lw'] = long_window
return stats
sws, lws = np.mgrid[50:80:5, 100:140:5]
#cProfile.run('run()')
stats_all = map(run, zip(sws.flatten(), lws.flatten()))
stats = run()
stats.returns.plot()
# for sw, lw in zip(sws.flatten(), lws.flatten()):
# stats = run(short_window=sw, long_window=lw)
# stats_all.append(stats)
stats = pd.concat(stats_all)
returns = stats.groupby(['sw', 'lw']).sum()
plt.contourf(sws, lws, returns.returns.reshape(sws.shape))
plt.xlabel('Short window length')
plt.ylabel('Long window length')
plt.savefig('DMA_contour.png')
plt.show()