mirror of
https://github.com/wassname/catalyst.git
synced 2026-07-23 12:50:22 +08:00
@@ -40,7 +40,7 @@ from zipline.finance.slippage import (
|
||||
)
|
||||
from zipline.finance.commission import PerShare, PerTrade
|
||||
from zipline.finance.blotter import Blotter
|
||||
from zipline.finance.constants import ANNUALIZER
|
||||
from zipline.finance.constants import ANNUALIZER, FILL_DELAYS
|
||||
import zipline.finance.trading as trading
|
||||
import zipline.protocol
|
||||
from zipline.protocol import Event
|
||||
@@ -87,6 +87,8 @@ class TradingAlgorithm(object):
|
||||
annualizer : int <optional>
|
||||
Which constant to use for annualizing risk metrics.
|
||||
If not provided, will extract from data_frequency.
|
||||
fill_delay : datetime.timedelta
|
||||
Delay between placing an order and filling an order.
|
||||
capital_base : float <default: 1.0e5>
|
||||
How much capital to start with.
|
||||
"""
|
||||
@@ -107,14 +109,13 @@ class TradingAlgorithm(object):
|
||||
self.slippage = VolumeShareSlippage()
|
||||
self.commission = PerShare()
|
||||
|
||||
if 'data_frequency' in kwargs:
|
||||
self.set_data_frequency(kwargs.pop('data_frequency'))
|
||||
else:
|
||||
self.data_frequency = None
|
||||
self.set_data_frequency(kwargs.pop('data_frequency', 'daily'))
|
||||
|
||||
# Override annualizer if set
|
||||
if 'annualizer' in kwargs:
|
||||
self.annualizer = kwargs['annualizer']
|
||||
if 'fill_delay' in kwargs:
|
||||
self.fill_delay = kwargs['fill_delay']
|
||||
|
||||
# set the capital base
|
||||
self.capital_base = kwargs.pop('capital_base', DEFAULT_CAPITAL_BASE)
|
||||
@@ -125,7 +126,7 @@ class TradingAlgorithm(object):
|
||||
|
||||
self.blotter = kwargs.pop('blotter', None)
|
||||
if not self.blotter:
|
||||
self.blotter = Blotter()
|
||||
self.blotter = Blotter(fill_delay=self.fill_delay)
|
||||
|
||||
# an algorithm subclass needs to set initialized to True when
|
||||
# it is fully initialized.
|
||||
@@ -431,3 +432,4 @@ class TradingAlgorithm(object):
|
||||
assert data_frequency in ('daily', 'minute')
|
||||
self.data_frequency = data_frequency
|
||||
self.annualizer = ANNUALIZER[self.data_frequency]
|
||||
self.fill_delay = FILL_DELAYS[self.data_frequency]
|
||||
|
||||
@@ -18,6 +18,7 @@ import uuid
|
||||
from copy import copy
|
||||
from logbook import Logger
|
||||
from collections import defaultdict
|
||||
from datetime import timedelta
|
||||
|
||||
import zipline.errors
|
||||
import zipline.protocol as zp
|
||||
@@ -43,7 +44,7 @@ ORDER_STATUS = Enum(
|
||||
|
||||
class Blotter(object):
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, fill_delay=timedelta(minutes=1)):
|
||||
self.transact = transact_partial(VolumeShareSlippage(), PerShare())
|
||||
# these orders are aggregated by sid
|
||||
self.open_orders = defaultdict(list)
|
||||
@@ -55,6 +56,8 @@ class Blotter(object):
|
||||
self.current_dt = None
|
||||
self.max_shares = int(1e+11)
|
||||
|
||||
self.fill_delay = fill_delay
|
||||
|
||||
def __repr__(self):
|
||||
return """
|
||||
{class_name}(
|
||||
@@ -155,8 +158,10 @@ class Blotter(object):
|
||||
orders = self.open_orders[trade_event.sid]
|
||||
orders = sorted(orders, key=lambda o: o.dt)
|
||||
# Only use orders for the current day or before
|
||||
# Since orders generally do not get filled immediately,
|
||||
# we allow for a delay here.
|
||||
current_orders = filter(
|
||||
lambda o: o.dt <= trade_event.dt,
|
||||
lambda o: o.dt + self.fill_delay <= trade_event.dt,
|
||||
orders)
|
||||
else:
|
||||
return
|
||||
|
||||
@@ -13,6 +13,8 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from datetime import timedelta
|
||||
|
||||
TRADING_DAYS_IN_YEAR = 250
|
||||
TRADING_HOURS_IN_DAY = 6
|
||||
MINUTES_IN_HOUR = 60
|
||||
@@ -21,3 +23,6 @@ ANNUALIZER = {'daily': TRADING_DAYS_IN_YEAR,
|
||||
'hourly': TRADING_DAYS_IN_YEAR * TRADING_HOURS_IN_DAY,
|
||||
'minute': TRADING_DAYS_IN_YEAR * TRADING_HOURS_IN_DAY *
|
||||
MINUTES_IN_HOUR}
|
||||
|
||||
FILL_DELAYS = {'daily': timedelta(days=1),
|
||||
'minute': timedelta(minutes=1)}
|
||||
|
||||
@@ -110,7 +110,7 @@ class AlgorithmSimulator(object):
|
||||
self.algo.perf_tracker.process_event(event)
|
||||
|
||||
else:
|
||||
|
||||
events = []
|
||||
for event in snapshot:
|
||||
if event.type in (DATASOURCE_TYPE.TRADE,
|
||||
DATASOURCE_TYPE.CUSTOM):
|
||||
@@ -120,11 +120,9 @@ class AlgorithmSimulator(object):
|
||||
self.algo.set_datetime(event.dt)
|
||||
bm_updated = True
|
||||
|
||||
process_trade = self.algo.blotter.process_trade
|
||||
for txn, order in process_trade(event):
|
||||
self.algo.perf_tracker.process_event(txn)
|
||||
self.algo.perf_tracker.process_event(order)
|
||||
self.algo.perf_tracker.process_event(event)
|
||||
# Save events to stream through blotter below.
|
||||
events.append(event)
|
||||
|
||||
|
||||
# Update our portfolio.
|
||||
self.algo.set_portfolio(
|
||||
@@ -145,6 +143,17 @@ class AlgorithmSimulator(object):
|
||||
self.algo.perf_tracker.process_event(order)
|
||||
self.algo.blotter.new_orders = []
|
||||
|
||||
# Fill orders
|
||||
for event in events:
|
||||
process_trade = self.algo.blotter.process_trade
|
||||
for txn, order in process_trade(event):
|
||||
|
||||
self.algo.perf_tracker.process_event(txn)
|
||||
self.algo.perf_tracker.process_event(order)
|
||||
|
||||
self.algo.perf_tracker.process_event(event)
|
||||
|
||||
|
||||
# The benchmark is our internal clock. When it
|
||||
# updates, we need to emit a performance message.
|
||||
if bm_updated:
|
||||
|
||||
Reference in New Issue
Block a user