mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-04 12:45:06 +08:00
ENH: Changed algorithm class to provide .run() method. Testing of multiple param combos of DMA.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user