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:
jfkirk
2015-09-10 11:53:28 -04:00
parent 661314ce49
commit dc964a7e7d
45 changed files with 1484 additions and 1173 deletions
+170 -141
View File
@@ -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']