mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-09 11:19:23 +08:00
MAINT: Removes the ability to reference a global TradingEnvironment
This commit removes the ability to reference a shared TradingEnvironment through the zipline.finance.trading module. In place, the classes that require a TradingEnvironment, or its child AssetFinder, contain their own references to those objects. This commit also adds serialization utilities that allow for the pickling/unpickling of objects without unintentionally their TradingEnvironments or AssetFinders.
This commit is contained in:
+170
-141
@@ -15,7 +15,6 @@
|
||||
|
||||
from __future__ import division
|
||||
|
||||
import pickle
|
||||
import collections
|
||||
from datetime import (
|
||||
datetime,
|
||||
@@ -40,12 +39,14 @@ from zipline.finance.slippage import Transaction, create_transaction
|
||||
import zipline.utils.math_utils as zp_math
|
||||
|
||||
from zipline.gens.composites import date_sorted_sources
|
||||
from zipline.finance import trading
|
||||
from zipline.finance.trading import SimulationParameters
|
||||
from zipline.finance.blotter import Order
|
||||
from zipline.finance.commission import PerShare, PerTrade, PerDollar
|
||||
from zipline.finance.trading import with_environment
|
||||
from zipline.finance.trading import TradingEnvironment
|
||||
from zipline.utils.factory import create_random_simulation_parameters
|
||||
from zipline.utils.serialization_utils import (
|
||||
load_with_persistent_ids, dump_with_persistent_ids
|
||||
)
|
||||
import zipline.protocol as zp
|
||||
from zipline.protocol import Event, DATASOURCE_TYPE
|
||||
from zipline.sources.data_frame_source import DataPanelSource
|
||||
@@ -128,8 +129,7 @@ def create_txn(trade_event, price, amount):
|
||||
return create_transaction(trade_event, mock_order, price, amount)
|
||||
|
||||
|
||||
@with_environment()
|
||||
def benchmark_events_in_range(sim_params, env=None):
|
||||
def benchmark_events_in_range(sim_params, env):
|
||||
return [
|
||||
Event({'dt': dt,
|
||||
'returns': ret,
|
||||
@@ -174,7 +174,7 @@ def calculate_results(host,
|
||||
txns = txns or []
|
||||
splits = splits or []
|
||||
|
||||
perf_tracker = perf.PerformanceTracker(host.sim_params)
|
||||
perf_tracker = perf.PerformanceTracker(host.sim_params, host.env)
|
||||
|
||||
if dividend_events is not None:
|
||||
dividend_frame = pd.DataFrame(
|
||||
@@ -246,9 +246,9 @@ def check_perf_tracker_serialization(perf_tracker):
|
||||
'total_days',
|
||||
]
|
||||
|
||||
p_string = pickle.dumps(perf_tracker)
|
||||
p_string = dump_with_persistent_ids(perf_tracker)
|
||||
|
||||
test = pickle.loads(p_string)
|
||||
test = load_with_persistent_ids(p_string, env=perf_tracker.env)
|
||||
|
||||
for k in scalar_keys:
|
||||
nt.assert_equal(getattr(test, k), getattr(perf_tracker, k), k)
|
||||
@@ -259,13 +259,15 @@ def check_perf_tracker_serialization(perf_tracker):
|
||||
|
||||
class TestSplitPerformance(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.env = TradingEnvironment()
|
||||
self.env.write_data(equities_identifiers=[1])
|
||||
self.sim_params, self.dt, self.end_dt = \
|
||||
create_random_simulation_parameters()
|
||||
trading.environment.write_data(equities_identifiers=[1])
|
||||
# start with $10,000
|
||||
self.sim_params.capital_base = 10e3
|
||||
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params)
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params,
|
||||
self.env)
|
||||
|
||||
def test_split_long_position(self):
|
||||
events = factory.create_trade_history(
|
||||
@@ -273,7 +275,8 @@ class TestSplitPerformance(unittest.TestCase):
|
||||
[20, 20],
|
||||
[100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
# set up a long position in sid 1
|
||||
@@ -359,17 +362,20 @@ class TestSplitPerformance(unittest.TestCase):
|
||||
class TestCommissionEvents(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
self.env = TradingEnvironment()
|
||||
self.env.write_data(
|
||||
equities_identifiers=[0, 1, 133]
|
||||
)
|
||||
self.sim_params, self.dt, self.end_dt = \
|
||||
create_random_simulation_parameters()
|
||||
|
||||
trading.environment.write_data(equities_identifiers=[0, 1, 133])
|
||||
|
||||
logger.info("sim_params: %s, dt: %s, end_dt: %s" %
|
||||
(self.sim_params, self.dt, self.end_dt))
|
||||
|
||||
self.sim_params.capital_base = 10e3
|
||||
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params)
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params,
|
||||
self.env)
|
||||
|
||||
def test_commission_event(self):
|
||||
events = factory.create_trade_history(
|
||||
@@ -377,7 +383,8 @@ class TestCommissionEvents(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
# Test commission models and validate result
|
||||
@@ -454,7 +461,8 @@ class TestCommissionEvents(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
# Buy and sell the same sid so that we have a zero position by the
|
||||
@@ -484,7 +492,8 @@ class TestCommissionEvents(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
# Add a cash adjustment at the time of event[3].
|
||||
@@ -500,21 +509,26 @@ class TestCommissionEvents(unittest.TestCase):
|
||||
|
||||
class TestDividendPerformance(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.env = TradingEnvironment()
|
||||
cls.env.write_data(equities_identifiers=[1, 2])
|
||||
|
||||
def setUp(self):
|
||||
self.sim_params, self.dt, self.end_dt = \
|
||||
create_random_simulation_parameters()
|
||||
trading.environment.write_data(equities_identifiers=[1, 2])
|
||||
self.sim_params.capital_base = 10e3
|
||||
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params)
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params,
|
||||
self.env)
|
||||
|
||||
def test_market_hours_calculations(self):
|
||||
# DST in US/Eastern began on Sunday March 14, 2010
|
||||
before = datetime(2010, 3, 12, 14, 31, tzinfo=pytz.utc)
|
||||
after = factory.get_next_trading_dt(
|
||||
before,
|
||||
timedelta(days=1)
|
||||
timedelta(days=1),
|
||||
self.env,
|
||||
)
|
||||
self.assertEqual(after.hour, 13)
|
||||
|
||||
@@ -525,7 +539,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
dividend = factory.create_dividend(
|
||||
1,
|
||||
@@ -576,7 +591,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params)
|
||||
self.sim_params,
|
||||
env=self.env)
|
||||
)
|
||||
|
||||
dividend = factory.create_stock_dividend(
|
||||
@@ -626,7 +642,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
dividend = factory.create_dividend(
|
||||
@@ -667,7 +684,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
dividend = factory.create_dividend(
|
||||
@@ -708,7 +726,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
dividend = factory.create_dividend(
|
||||
@@ -749,13 +768,14 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
pay_date = self.sim_params.first_open
|
||||
# find pay date that is much later.
|
||||
for i in range(30):
|
||||
pay_date = factory.get_next_trading_dt(pay_date, oneday)
|
||||
pay_date = factory.get_next_trading_dt(pay_date, oneday, self.env)
|
||||
dividend = factory.create_dividend(
|
||||
1,
|
||||
10.00,
|
||||
@@ -795,7 +815,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
dividend = factory.create_dividend(
|
||||
@@ -836,7 +857,8 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
dividend = factory.create_dividend(
|
||||
@@ -865,15 +887,15 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
[event['cumulative_perf']['capital_used'] for event in results]
|
||||
self.assertEqual(cumulative_cash_flows, [0, 0, 0, 0, 0])
|
||||
|
||||
@with_environment()
|
||||
def test_no_dividend_at_simulation_end(self, env=None):
|
||||
def test_no_dividend_at_simulation_end(self):
|
||||
# post some trades in the market
|
||||
events = factory.create_trade_history(
|
||||
1,
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
dividend = factory.create_dividend(
|
||||
1,
|
||||
@@ -886,12 +908,12 @@ class TestDividendPerformance(unittest.TestCase):
|
||||
events[-2].dt,
|
||||
# pay date, when the algorithm receives the dividend.
|
||||
# This pays out on the day after the last event
|
||||
env.next_trading_day(events[-1].dt)
|
||||
self.env.next_trading_day(events[-1].dt)
|
||||
)
|
||||
|
||||
# Set the last day to be the last event
|
||||
self.sim_params.period_end = events[-1].dt
|
||||
self.sim_params._update_internal()
|
||||
self.sim_params.update_internal_from_env(self.env)
|
||||
|
||||
# Simulate a transaction being filled prior to the ex_date.
|
||||
txns = [create_txn(events[0], 10.0, 100)]
|
||||
@@ -929,18 +951,29 @@ class TestDividendPerformanceHolidayStyle(TestDividendPerformance):
|
||||
self.end_dt = datetime(2004, 11, 25, tzinfo=pytz.utc)
|
||||
self.sim_params = SimulationParameters(
|
||||
self.dt,
|
||||
self.end_dt)
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params)
|
||||
self.end_dt,
|
||||
env=self.env)
|
||||
|
||||
self.sim_params.capital_base = 10e3
|
||||
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params,
|
||||
self.env)
|
||||
|
||||
|
||||
class TestPositionPerformance(unittest.TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.env = TradingEnvironment()
|
||||
cls.env.write_data(equities_identifiers=[1, 2])
|
||||
|
||||
def setUp(self):
|
||||
self.sim_params, self.dt, self.end_dt = \
|
||||
create_random_simulation_parameters()
|
||||
|
||||
trading.environment.write_data(equities_identifiers=[1, 2])
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params)
|
||||
self.finder = self.env.asset_finder
|
||||
self.benchmark_events = benchmark_events_in_range(self.sim_params,
|
||||
self.env)
|
||||
|
||||
def test_long_short_positions(self):
|
||||
"""
|
||||
@@ -956,7 +989,8 @@ class TestPositionPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 9],
|
||||
[100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
trades_2 = factory.create_trade_history(
|
||||
@@ -964,13 +998,14 @@ class TestPositionPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 11],
|
||||
[100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
txn1 = create_txn(trades_1[1], 10.0, 100)
|
||||
txn2 = create_txn(trades_2[1], 10.0, -100)
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
pt.execute_transaction(txn1)
|
||||
pp.handle_execution(txn1)
|
||||
@@ -1046,12 +1081,13 @@ class TestPositionPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 11],
|
||||
[100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
txn = create_txn(trades[1], 10.0, 1000)
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
pt.execute_transaction(txn)
|
||||
@@ -1125,12 +1161,13 @@ class TestPositionPerformance(unittest.TestCase):
|
||||
[10, 10, 10, 11],
|
||||
[100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
txn = create_txn(trades[1], 10.0, 100)
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
pt.execute_transaction(txn)
|
||||
@@ -1228,14 +1265,15 @@ single short-sale transaction"""
|
||||
[10, 10, 10, 11, 10, 9],
|
||||
[100, 100, 100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
trades_1 = trades[:-2]
|
||||
|
||||
txn = create_txn(trades[1], 10.0, -100)
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
pt.execute_transaction(txn)
|
||||
@@ -1352,8 +1390,8 @@ single short-sale transaction"""
|
||||
)
|
||||
|
||||
# now run a performance period encompassing the entire trade sample.
|
||||
ptTotal = perf.PositionTracker()
|
||||
ppTotal = perf.PerformancePeriod(1000.0)
|
||||
ptTotal = perf.PositionTracker(self.env.asset_finder)
|
||||
ppTotal = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
ppTotal.position_tracker = pt
|
||||
|
||||
for trade in trades_1:
|
||||
@@ -1447,7 +1485,8 @@ trade after cover"""
|
||||
[10, 10, 10, 11, 9, 8, 7, 8, 9, 10],
|
||||
[100, 100, 100, 100, 100, 100, 100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
short_txn = create_txn(
|
||||
@@ -1457,8 +1496,8 @@ trade after cover"""
|
||||
)
|
||||
|
||||
cover_txn = create_txn(trades[6], 7.0, 100)
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
pt.execute_transaction(short_txn)
|
||||
@@ -1551,13 +1590,14 @@ shares in position"
|
||||
[10, 11, 11, 12],
|
||||
[100, 100, 100, 100],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
self.env
|
||||
)
|
||||
trades = factory.create_trade_history(*history_args)
|
||||
transactions = factory.create_txn_history(*history_args)
|
||||
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
average_cost = 0
|
||||
@@ -1623,8 +1663,8 @@ shares in position"
|
||||
|
||||
self.assertEqual(pp.pnl, -800, "this period goes from +400 to -400")
|
||||
|
||||
pt3 = perf.PositionTracker()
|
||||
pp3 = perf.PerformancePeriod(1000.0)
|
||||
pt3 = perf.PositionTracker(self.env.asset_finder)
|
||||
pp3 = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp3.position_tracker = pt3
|
||||
|
||||
average_cost = 0
|
||||
@@ -1666,15 +1706,16 @@ shares in position"
|
||||
[10, 9, 11, 8, 9, 12, 13, 14],
|
||||
[200, -100, -100, 100, -300, 100, 500, 400],
|
||||
onesec,
|
||||
self.sim_params
|
||||
self.sim_params,
|
||||
self.env
|
||||
)
|
||||
cost_bases = [10, 10, 0, 8, 9, 9, 13, 13.5]
|
||||
|
||||
trades = factory.create_trade_history(*history_args)
|
||||
transactions = factory.create_txn_history(*history_args)
|
||||
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(1000.0)
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
pp = perf.PerformancePeriod(1000.0, self.env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
for txn, cb in zip(transactions, cost_bases):
|
||||
@@ -1692,9 +1733,10 @@ shares in position"
|
||||
|
||||
class TestPerformanceTracker(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
trading.environment = trading.TradingEnvironment()
|
||||
trading.environment.write_data(equities_identifiers=[133, 134])
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.env = TradingEnvironment()
|
||||
cls.env.write_data(equities_identifiers=[1, 2, 133, 134])
|
||||
|
||||
NumDaysToDelete = collections.namedtuple(
|
||||
'NumDaysToDelete', ('start', 'middle', 'end'))
|
||||
@@ -1733,8 +1775,6 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
# 12 13 14 15 16 17 18
|
||||
# 19 20 21 22 23 24 25
|
||||
# 26 27 28 29 30 31
|
||||
trading.environment = trading.TradingEnvironment()
|
||||
trading.environment.write_data(equities_identifiers=[133, 134])
|
||||
start_dt = datetime(year=2008,
|
||||
month=10,
|
||||
day=9,
|
||||
@@ -1753,10 +1793,11 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
|
||||
sim_params = SimulationParameters(
|
||||
period_start=start_dt,
|
||||
period_end=end_dt
|
||||
period_end=end_dt,
|
||||
env=self.env,
|
||||
)
|
||||
|
||||
benchmark_events = benchmark_events_in_range(sim_params)
|
||||
benchmark_events = benchmark_events_in_range(sim_params, self.env)
|
||||
|
||||
trade_history = factory.create_trade_history(
|
||||
sid,
|
||||
@@ -1764,7 +1805,8 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
volume,
|
||||
trade_time_increment,
|
||||
sim_params,
|
||||
source_id="factory1"
|
||||
source_id="factory1",
|
||||
env=self.env
|
||||
)
|
||||
|
||||
sid2 = 134
|
||||
@@ -1776,7 +1818,8 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
volume,
|
||||
trade_time_increment,
|
||||
sim_params,
|
||||
source_id="factory2"
|
||||
source_id="factory2",
|
||||
env=self.env
|
||||
)
|
||||
# 'middle' start of 3 depends on number of days == 7
|
||||
middle = 3
|
||||
@@ -1796,10 +1839,6 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
del trade_history[-days_to_delete.end:]
|
||||
del trade_history2[-days_to_delete.end:]
|
||||
|
||||
sim_params.first_open = \
|
||||
sim_params.calculate_first_open()
|
||||
sim_params.last_close = \
|
||||
sim_params.calculate_last_close()
|
||||
sim_params.capital_base = 1000.0
|
||||
sim_params.frame_index = [
|
||||
'sid',
|
||||
@@ -1808,7 +1847,7 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
'price',
|
||||
'changed']
|
||||
perf_tracker = perf.PerformanceTracker(
|
||||
sim_params
|
||||
sim_params, self.env
|
||||
)
|
||||
|
||||
events = date_sorted_sources(trade_history, trade_history2)
|
||||
@@ -1887,23 +1926,21 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
else:
|
||||
yield event
|
||||
|
||||
@with_environment()
|
||||
def test_minute_tracker(self, env=None):
|
||||
def test_minute_tracker(self):
|
||||
""" Tests minute performance tracking."""
|
||||
start_dt = env.exchange_dt_in_utc(datetime(2013, 3, 1, 9, 31))
|
||||
end_dt = env.exchange_dt_in_utc(datetime(2013, 3, 1, 16, 0))
|
||||
|
||||
sim_params = SimulationParameters(
|
||||
period_start=start_dt,
|
||||
period_end=end_dt,
|
||||
emission_rate='minute'
|
||||
)
|
||||
tracker = perf.PerformanceTracker(sim_params)
|
||||
start_dt = self.env.exchange_dt_in_utc(datetime(2013, 3, 1, 9, 31))
|
||||
end_dt = self.env.exchange_dt_in_utc(datetime(2013, 3, 1, 16, 0))
|
||||
|
||||
foosid = 1
|
||||
barsid = 2
|
||||
|
||||
env.write_data(equities_identifiers=[foosid, barsid])
|
||||
sim_params = SimulationParameters(
|
||||
period_start=start_dt,
|
||||
period_end=end_dt,
|
||||
emission_rate='minute',
|
||||
env=self.env,
|
||||
)
|
||||
tracker = perf.PerformanceTracker(sim_params, env=self.env)
|
||||
|
||||
foo_event_1 = factory.create_trade(foosid, 10.0, 20, start_dt)
|
||||
order_event_1 = Order(sid=foo_event_1.sid,
|
||||
@@ -1996,10 +2033,8 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
|
||||
check_perf_tracker_serialization(tracker)
|
||||
|
||||
@with_environment()
|
||||
def test_close_position_event(self, env=None):
|
||||
env.write_data(equities_identifiers=[1, 2])
|
||||
pt = perf.PositionTracker()
|
||||
def test_close_position_event(self):
|
||||
pt = perf.PositionTracker(asset_finder=self.env.asset_finder)
|
||||
dt = pd.Timestamp("1984/03/06 3:00PM")
|
||||
pos1 = perf.Position(1, amount=np.float64(120.0),
|
||||
last_sale_date=dt, last_sale_price=3.4)
|
||||
@@ -2037,11 +2072,12 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
[10, 10, 10, 10, 10],
|
||||
[100, 100, 100, 100, 100],
|
||||
oneday,
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
|
||||
# Create a tracker and a dividend
|
||||
perf_tracker = perf.PerformanceTracker(sim_params)
|
||||
perf_tracker = perf.PerformanceTracker(sim_params, env=self.env)
|
||||
dividend = factory.create_dividend(
|
||||
1,
|
||||
10.00,
|
||||
@@ -2081,11 +2117,12 @@ class TestPerformanceTracker(unittest.TestCase):
|
||||
|
||||
sim_params = SimulationParameters(
|
||||
period_start=start_dt,
|
||||
period_end=end_dt
|
||||
period_end=end_dt,
|
||||
env=self.env,
|
||||
)
|
||||
|
||||
perf_tracker = perf.PerformanceTracker(
|
||||
sim_params
|
||||
sim_params, env=self.env
|
||||
)
|
||||
check_perf_tracker_serialization(perf_tracker)
|
||||
|
||||
@@ -2099,16 +2136,26 @@ class TestPosition(unittest.TestCase):
|
||||
pos = perf.Position(10, amount=np.float64(120.0), last_sale_date=dt,
|
||||
last_sale_price=3.4)
|
||||
|
||||
p_string = pickle.dumps(pos)
|
||||
p_string = dump_with_persistent_ids(pos)
|
||||
|
||||
test = pickle.loads(p_string)
|
||||
test = load_with_persistent_ids(p_string, env=None)
|
||||
nt.assert_dict_equal(test.__dict__, pos.__dict__)
|
||||
|
||||
|
||||
class TestPositionTracker(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
trading.environment = trading.TradingEnvironment()
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.env = TradingEnvironment()
|
||||
|
||||
equities_metadata = {1: {'asset_type': 'equity'},
|
||||
2: {'asset_type': 'equity'}}
|
||||
futures_metadata = {3: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000},
|
||||
4: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000}}
|
||||
cls.env.write_data(equities_data=equities_metadata,
|
||||
futures_data=futures_metadata)
|
||||
|
||||
def test_empty_positions(self):
|
||||
"""
|
||||
@@ -2117,7 +2164,7 @@ class TestPositionTracker(unittest.TestCase):
|
||||
Originally this bug was due to np.dot([], []) returning
|
||||
np.bool_(False)
|
||||
"""
|
||||
pt = perf.PositionTracker()
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
|
||||
stats = [
|
||||
'calculate_positions_value',
|
||||
@@ -2137,41 +2184,28 @@ class TestPositionTracker(unittest.TestCase):
|
||||
self.assertEquals(val, 0)
|
||||
self.assertNotIsInstance(val, (bool, np.bool_))
|
||||
|
||||
def test_update_last_sale(self, env=None):
|
||||
equities_metadata = {1: {'asset_type': 'equity'}}
|
||||
futures_metadata = {2: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000}}
|
||||
trading.environment.write_data(equities_data=equities_metadata,
|
||||
futures_data=futures_metadata)
|
||||
pt = perf.PositionTracker()
|
||||
def test_update_last_sale(self):
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
dt = pd.Timestamp("1984/03/06 3:00PM")
|
||||
pos1 = perf.Position(1, amount=np.float64(100.0),
|
||||
last_sale_date=dt, last_sale_price=10)
|
||||
pos2 = perf.Position(2, amount=np.float64(100.0),
|
||||
pos3 = perf.Position(3, amount=np.float64(100.0),
|
||||
last_sale_date=dt, last_sale_price=10)
|
||||
pt.update_positions({1: pos1, 2: pos2})
|
||||
pt.update_positions({1: pos1, 3: pos3})
|
||||
|
||||
event1 = Event({'sid': 1,
|
||||
'price': 11,
|
||||
'dt': dt})
|
||||
event2 = Event({'sid': 2,
|
||||
event3 = Event({'sid': 3,
|
||||
'price': 11,
|
||||
'dt': dt})
|
||||
|
||||
# Check cash-adjustment return value
|
||||
self.assertEqual(0, pt.update_last_sale(event1))
|
||||
self.assertEqual(100000, pt.update_last_sale(event2))
|
||||
self.assertEqual(100000, pt.update_last_sale(event3))
|
||||
|
||||
def test_position_values_and_exposures(self, env=None):
|
||||
equities_metadata = {1: {'asset_type': 'equity'},
|
||||
2: {'asset_type': 'equity'}}
|
||||
futures_metadata = {3: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000},
|
||||
4: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000}}
|
||||
trading.environment.write_data(equities_data=equities_metadata,
|
||||
futures_data=futures_metadata)
|
||||
pt = perf.PositionTracker()
|
||||
def test_position_values_and_exposures(self):
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
dt = pd.Timestamp("1984/03/06 3:00PM")
|
||||
pos1 = perf.Position(1, amount=np.float64(10.0),
|
||||
last_sale_date=dt, last_sale_price=10)
|
||||
@@ -2199,21 +2233,17 @@ class TestPositionTracker(unittest.TestCase):
|
||||
self.assertEqual(100 + 200 + 300000 + 400000, pt._gross_exposure())
|
||||
self.assertEqual(100 - 200 + 300000 - 400000, pt._net_exposure())
|
||||
|
||||
def test_serialization(self, env=None):
|
||||
metadata = {1: {'asset_type': 'equity'},
|
||||
2: {'asset_type': 'future',
|
||||
'contract_multiplier': 1000}}
|
||||
trading.environment.write_data(equities_data=metadata)
|
||||
pt = perf.PositionTracker()
|
||||
def test_serialization(self):
|
||||
pt = perf.PositionTracker(self.env.asset_finder)
|
||||
dt = pd.Timestamp("1984/03/06 3:00PM")
|
||||
pos1 = perf.Position(1, amount=np.float64(120.0),
|
||||
last_sale_date=dt, last_sale_price=3.4)
|
||||
pos2 = perf.Position(2, amount=np.float64(100.0),
|
||||
pos3 = perf.Position(3, amount=np.float64(100.0),
|
||||
last_sale_date=dt, last_sale_price=3.4)
|
||||
|
||||
pt.update_positions({1: pos1, 2: pos2})
|
||||
p_string = pickle.dumps(pt)
|
||||
test = pickle.loads(p_string)
|
||||
pt.update_positions({1: pos1, 3: pos3})
|
||||
p_string = dump_with_persistent_ids(pt)
|
||||
test = load_with_persistent_ids(p_string, env=self.env)
|
||||
nt.assert_dict_equal(test._position_amounts, pt._position_amounts)
|
||||
nt.assert_dict_equal(test._position_last_sale_prices,
|
||||
pt._position_last_sale_prices)
|
||||
@@ -2224,16 +2254,15 @@ class TestPositionTracker(unittest.TestCase):
|
||||
|
||||
|
||||
class TestPerformancePeriod(unittest.TestCase):
|
||||
def setUp(self):
|
||||
pass
|
||||
|
||||
def test_serialization(self):
|
||||
pt = perf.PositionTracker()
|
||||
pp = perf.PerformancePeriod(100)
|
||||
env = TradingEnvironment()
|
||||
pt = perf.PositionTracker(env.asset_finder)
|
||||
pp = perf.PerformancePeriod(100, env.asset_finder)
|
||||
pp.position_tracker = pt
|
||||
|
||||
p_string = pickle.dumps(pp)
|
||||
test = pickle.loads(p_string)
|
||||
p_string = dump_with_persistent_ids(pp)
|
||||
test = load_with_persistent_ids(p_string, env=env)
|
||||
|
||||
correct = pp.__dict__.copy()
|
||||
del correct['_position_tracker']
|
||||
|
||||
Reference in New Issue
Block a user