mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-09 11:19:23 +08:00
Merge pull request #1313 from nathanwolfe/master
BUG: Add support for Panel data in accordance with documentation
This commit is contained in:
+85
-1
@@ -33,7 +33,10 @@ import numpy as np
|
||||
import pandas as pd
|
||||
import pytz
|
||||
|
||||
from zipline import TradingAlgorithm
|
||||
from zipline import (
|
||||
run_algorithm,
|
||||
TradingAlgorithm,
|
||||
)
|
||||
from zipline.api import FixedSlippage
|
||||
from zipline.assets import Equity, Future
|
||||
from zipline.assets.synthetic import (
|
||||
@@ -161,6 +164,7 @@ from zipline.test_algorithms import (
|
||||
no_handle_data,
|
||||
)
|
||||
from zipline.utils.api_support import ZiplineAPI, set_algo_instance
|
||||
from zipline.utils.calendars import get_calendar
|
||||
from zipline.utils.context_tricks import CallbackManager
|
||||
from zipline.utils.control_flow import nullctx
|
||||
import zipline.utils.events
|
||||
@@ -4102,3 +4106,83 @@ class AlgoInputValidationTestCase(ZiplineTestCase):
|
||||
script=script,
|
||||
**{method: lambda *args, **kwargs: None}
|
||||
)
|
||||
|
||||
|
||||
class TestPanelData(ZiplineTestCase):
|
||||
|
||||
@parameterized.expand([
|
||||
('daily',
|
||||
pd.Timestamp('2015-12-23', tz='UTC'),
|
||||
pd.Timestamp('2016-01-05', tz='UTC'),),
|
||||
('minute',
|
||||
pd.Timestamp('2015-12-23', tz='UTC'),
|
||||
pd.Timestamp('2015-12-24', tz='UTC'),),
|
||||
])
|
||||
def test_panel_data(self, data_frequency, start_dt, end_dt):
|
||||
trading_calendar = get_calendar('NYSE')
|
||||
if data_frequency == 'daily':
|
||||
history_freq = '1d'
|
||||
create_df_for_asset = create_daily_df_for_asset
|
||||
dt_transform = trading_calendar.minute_to_session_label
|
||||
elif data_frequency == 'minute':
|
||||
history_freq = '1m'
|
||||
create_df_for_asset = create_minute_df_for_asset
|
||||
|
||||
def dt_transform(dt):
|
||||
return dt
|
||||
|
||||
sids = range(1, 3)
|
||||
dfs = {}
|
||||
for sid in sids:
|
||||
dfs[sid] = create_df_for_asset(trading_calendar,
|
||||
start_dt, end_dt, interval=sid)
|
||||
dfs[sid]['prev_close'] = dfs[sid]['close'].shift(1)
|
||||
panel = pd.Panel(dfs)
|
||||
|
||||
price_record = pd.Panel(items=sids,
|
||||
major_axis=panel.major_axis,
|
||||
minor_axis=['current', 'previous'])
|
||||
|
||||
def initialize(algo):
|
||||
algo.first_bar = True
|
||||
algo.equities = []
|
||||
for sid in sids:
|
||||
algo.equities.append(algo.sid(sid))
|
||||
|
||||
def handle_data(algo, data):
|
||||
price_record.loc[:, dt_transform(algo.get_datetime()),
|
||||
'current'] = (
|
||||
data.current(algo.equities, 'price')
|
||||
)
|
||||
if algo.first_bar:
|
||||
algo.first_bar = False
|
||||
else:
|
||||
price_record.loc[:, dt_transform(algo.get_datetime()),
|
||||
'previous'] = (
|
||||
data.history(algo.equities, 'price',
|
||||
2, history_freq).iloc[0]
|
||||
)
|
||||
|
||||
def check_panels():
|
||||
np.testing.assert_array_equal(
|
||||
price_record.values.astype('float64'),
|
||||
panel.loc[:, :, ['close',
|
||||
'prev_close']].values.astype('float64')
|
||||
)
|
||||
|
||||
trading_algo = TradingAlgorithm(initialize=initialize,
|
||||
handle_data=handle_data)
|
||||
trading_algo.run(data=panel)
|
||||
check_panels()
|
||||
price_record.loc[:] = np.nan
|
||||
|
||||
run_algorithm(
|
||||
start=start_dt,
|
||||
end=end_dt,
|
||||
capital_base=1,
|
||||
initialize=initialize,
|
||||
handle_data=handle_data,
|
||||
data_frequency=data_frequency,
|
||||
data=panel
|
||||
)
|
||||
check_panels()
|
||||
|
||||
@@ -18,31 +18,29 @@ from itertools import permutations, product
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from zipline.data.us_equity_pricing import PanelDailyBarReader
|
||||
from zipline.data.us_equity_pricing import PanelBarReader
|
||||
from zipline.testing import ExplodingObject
|
||||
from zipline.testing.fixtures import (
|
||||
WithAssetFinder,
|
||||
WithNYSETradingDays,
|
||||
ZiplineTestCase,
|
||||
)
|
||||
from zipline.utils.calendars import get_calendar
|
||||
|
||||
|
||||
class TestPanelDailyBarReader(WithAssetFinder,
|
||||
WithNYSETradingDays,
|
||||
ZiplineTestCase):
|
||||
|
||||
START_DATE = pd.Timestamp('2006-01-03', tz='utc')
|
||||
END_DATE = pd.Timestamp('2006-02-01', tz='utc')
|
||||
class WithPanelBarReader(WithAssetFinder):
|
||||
|
||||
@classmethod
|
||||
def init_class_fixtures(cls):
|
||||
super(TestPanelDailyBarReader, cls).init_class_fixtures()
|
||||
super(WithPanelBarReader, cls).init_class_fixtures()
|
||||
|
||||
finder = cls.asset_finder
|
||||
days = cls.trading_days
|
||||
trading_calendar = get_calendar('NYSE')
|
||||
|
||||
items = finder.retrieve_all(finder.sids)
|
||||
major_axis = days
|
||||
major_axis = (
|
||||
trading_calendar.sessions_in_range if cls.FREQUENCY == 'daily'
|
||||
else trading_calendar.minutes_for_sessions_in_range
|
||||
)(cls.START_DATE, cls.END_DATE)
|
||||
minor_axis = ['open', 'high', 'low', 'close', 'volume']
|
||||
|
||||
shape = tuple(map(len, [items, major_axis, minor_axis]))
|
||||
@@ -55,7 +53,7 @@ class TestPanelDailyBarReader(WithAssetFinder,
|
||||
minor_axis=minor_axis,
|
||||
)
|
||||
|
||||
cls.reader = PanelDailyBarReader(days, cls.panel)
|
||||
cls.reader = PanelBarReader(trading_calendar, cls.panel, cls.FREQUENCY)
|
||||
|
||||
def test_spot_price(self):
|
||||
panel = self.panel
|
||||
@@ -83,7 +81,7 @@ class TestPanelDailyBarReader(WithAssetFinder,
|
||||
for axis_order in permutations((0, 1, 2)):
|
||||
transposed = panel.transpose(*axis_order)
|
||||
with self.assertRaises(ValueError) as e:
|
||||
PanelDailyBarReader(unused, transposed)
|
||||
PanelBarReader(unused, transposed, 'daily')
|
||||
|
||||
expected = (
|
||||
"Duplicate entries in Panel.{name}: ['a', 'b'].".format(
|
||||
@@ -95,6 +93,28 @@ class TestPanelDailyBarReader(WithAssetFinder,
|
||||
def test_sessions(self):
|
||||
sessions = self.reader.sessions
|
||||
|
||||
self.assertEqual(21, len(sessions))
|
||||
self.assertEqual(self.NUM_SESSIONS, len(sessions))
|
||||
self.assertEqual(self.START_DATE, sessions[0])
|
||||
self.assertEqual(self.END_DATE, sessions[-1])
|
||||
|
||||
|
||||
class TestPanelDailyBarReader(WithPanelBarReader,
|
||||
ZiplineTestCase):
|
||||
|
||||
FREQUENCY = 'daily'
|
||||
|
||||
START_DATE = pd.Timestamp('2006-01-03', tz='utc')
|
||||
END_DATE = pd.Timestamp('2006-02-01', tz='utc')
|
||||
|
||||
NUM_SESSIONS = 21
|
||||
|
||||
|
||||
class TestPanelMinuteBarReader(WithPanelBarReader,
|
||||
ZiplineTestCase):
|
||||
|
||||
FREQUENCY = 'minute'
|
||||
|
||||
START_DATE = pd.Timestamp('2015-12-23', tz='utc')
|
||||
END_DATE = pd.Timestamp('2015-12-24', tz='utc')
|
||||
|
||||
NUM_SESSIONS = 2
|
||||
Reference in New Issue
Block a user