mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-18 11:50:11 +08:00
ENH: Implemented new Zipline interface. Implemented Algorithm base class. Implemented example algorithm. Implemented example.py code.
This commit is contained in:
@@ -1,3 +1,6 @@
|
||||
from zipline.gens.mavg import MovingAverage
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
class BuySellAlgorithm(object):
|
||||
"""Algorithm that buys and sells alternatingly. The amount for
|
||||
each order can be specified. In addition, an offset that will
|
||||
@@ -46,3 +49,57 @@ class BuySellAlgorithm(object):
|
||||
|
||||
def get_sid_filter(self):
|
||||
return [self.sid]
|
||||
|
||||
# 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):
|
||||
def set_order(self, order_callable):
|
||||
self.order = order_callable
|
||||
|
||||
def get_sid_filter(self):
|
||||
return [self.sid]
|
||||
|
||||
def set_logger(self, logger):
|
||||
self.logger = logger
|
||||
|
||||
def initialize(self):
|
||||
pass
|
||||
|
||||
def add_transform(self, transform_class, tag, *args, **kwargs):
|
||||
if not hasattr(self, 'registered_transforms'):
|
||||
self.registered_transforms = {}
|
||||
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user