mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-12 11:50:11 +08:00
Refinements and documentation.
This commit is contained in:
@@ -29,7 +29,7 @@ from catalyst.exchange.exchange_errors import (
|
||||
ExchangeRequestError,
|
||||
ExchangeBarDataError,
|
||||
BundleNotFoundError, PricingDataBeforeTradingError,
|
||||
PricingDataNotLoadedError)
|
||||
PricingDataNotLoadedError, InvalidHistoryFrequencyError)
|
||||
from catalyst.utils.paths import data_path
|
||||
|
||||
log = Logger('DataPortalExchange')
|
||||
@@ -154,8 +154,9 @@ class DataPortalExchangeBase(DataPortal):
|
||||
try:
|
||||
if isinstance(assets, TradingPair):
|
||||
exchange = self.exchanges[assets.exchange]
|
||||
return self.get_exchange_spot_value(
|
||||
exchange, assets, field, dt, data_frequency)
|
||||
spot_values = self.get_exchange_spot_value(
|
||||
exchange, [assets], field, dt, data_frequency)
|
||||
return spot_values[0]
|
||||
|
||||
else:
|
||||
exchange_assets = dict()
|
||||
@@ -306,6 +307,15 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase):
|
||||
|
||||
@staticmethod
|
||||
def find_most_recent_time(bundle_name):
|
||||
"""
|
||||
Find most recent "time folder" for a given bundle.
|
||||
|
||||
:param bundle_name:
|
||||
The name of the targeted bundle.
|
||||
|
||||
:return folder:
|
||||
The name of the time folder.
|
||||
"""
|
||||
try:
|
||||
bundle_folders = os.listdir(
|
||||
data_path([bundle_name]),
|
||||
@@ -326,6 +336,33 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase):
|
||||
else:
|
||||
return None
|
||||
|
||||
def _get_reader(self, data_frequency, exchange_name):
|
||||
"""
|
||||
Pick from a collection of readers based of exchange name and frequency.
|
||||
|
||||
:param data_frequency:
|
||||
The reader frequency: minute, 5-minute, daily.
|
||||
|
||||
:param exchange_name:
|
||||
The exchange name.
|
||||
|
||||
:return reader:
|
||||
A reader object.
|
||||
"""
|
||||
if data_frequency == 'minute':
|
||||
reader = self.minute_bar_readers[exchange_name]
|
||||
elif data_frequency == '5-minute':
|
||||
reader = self.five_minute_bar_readers[exchange_name]
|
||||
elif data_frequency == 'daily':
|
||||
reader = self.daily_bar_readers[exchange_name]
|
||||
else:
|
||||
raise InvalidHistoryFrequencyError(frequency=data_frequency)
|
||||
|
||||
if reader is None:
|
||||
raise ValueError('reader not found')
|
||||
|
||||
return reader
|
||||
|
||||
def get_exchange_history_window(self,
|
||||
exchange,
|
||||
assets,
|
||||
@@ -335,33 +372,31 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase):
|
||||
field,
|
||||
data_frequency,
|
||||
ffill=True):
|
||||
if data_frequency == 'minute' or data_frequency == '1m':
|
||||
reader = self.minute_bar_readers[exchange.name]
|
||||
reader = self._get_reader(data_frequency, exchange.name)
|
||||
if data_frequency == 'minute':
|
||||
dts = self.trading_calendar.minutes_window(
|
||||
end_dt, -bar_count
|
||||
)
|
||||
|
||||
self.ensure_after_first_day(dts[0], assets)
|
||||
|
||||
elif data_frequency == '5-minute' or data_frequency == '5m':
|
||||
reader = self.five_minute_bar_readers[exchange.name]
|
||||
elif data_frequency == 'daily' or data_frequency == '1d':
|
||||
reader = self.daily_bar_readers[exchange.name]
|
||||
|
||||
elif data_frequency == 'daily':
|
||||
session = self.trading_calendar.minute_to_session_label(end_dt)
|
||||
dts = self._get_days_for_window(session, bar_count)
|
||||
|
||||
self.ensure_after_first_day(dts[0], assets)
|
||||
|
||||
else:
|
||||
raise ValueError('Unsupported frequency')
|
||||
raise InvalidHistoryFrequencyError(frequency=data_frequency)
|
||||
|
||||
try:
|
||||
values = reader.load_raw_arrays(
|
||||
[field],
|
||||
dts[0],
|
||||
dts[-1],
|
||||
assets,
|
||||
fields=[field],
|
||||
start_dt=dts[0],
|
||||
end_dt=dts[-1],
|
||||
sids=[asset.sid for asset in assets],
|
||||
)[0]
|
||||
|
||||
except Exception:
|
||||
raise PricingDataNotLoadedError(
|
||||
field=field,
|
||||
@@ -393,50 +428,25 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase):
|
||||
|
||||
def get_exchange_spot_value(self, exchange, assets, field, dt,
|
||||
data_frequency):
|
||||
if data_frequency == 'minute' or data_frequency == '1m':
|
||||
reader = self.minute_bar_readers[exchange.name]
|
||||
elif data_frequency == '5-minute' or data_frequency == '5m':
|
||||
reader = self.five_minute_bar_readers[exchange.name]
|
||||
elif data_frequency == 'daily' or data_frequency == '1d':
|
||||
reader = self.daily_bar_readers[exchange.name]
|
||||
else:
|
||||
raise ValueError('Unsupported frequency')
|
||||
reader = self._get_reader(data_frequency, exchange.name)
|
||||
|
||||
if isinstance(assets, TradingPair):
|
||||
self.ensure_after_first_day(dt, [assets])
|
||||
self.ensure_after_first_day(dt, assets)
|
||||
|
||||
values = []
|
||||
for asset in assets:
|
||||
try:
|
||||
value = reader.get_value(
|
||||
sid=assets.sid,
|
||||
sid=asset.sid,
|
||||
dt=dt,
|
||||
field=field
|
||||
)
|
||||
return value
|
||||
values.append(value)
|
||||
except Exception:
|
||||
raise PricingDataNotLoadedError(
|
||||
field=field,
|
||||
first_trading_day=self._get_first_trading_day([assets]),
|
||||
first_trading_day=self._get_first_trading_day(assets),
|
||||
exchange=exchange.name,
|
||||
symbols=assets.symbol,
|
||||
symbols=[asset.symbol for asset in assets],
|
||||
)
|
||||
else:
|
||||
self.ensure_after_first_day(dt, assets)
|
||||
|
||||
values = []
|
||||
for asset in assets:
|
||||
try:
|
||||
value = reader.get_value(
|
||||
sid=asset.sid,
|
||||
dt=dt,
|
||||
field=field
|
||||
)
|
||||
values.append(value)
|
||||
except Exception:
|
||||
raise PricingDataNotLoadedError(
|
||||
field=field,
|
||||
first_trading_day=self._get_first_trading_day(assets),
|
||||
exchange=exchange.name,
|
||||
symbols=[asset.symbol for asset in assets],
|
||||
)
|
||||
|
||||
return values
|
||||
return values
|
||||
|
||||
@@ -306,19 +306,14 @@ class Exchange:
|
||||
'1D', '7D', '14D', '1M'
|
||||
"""
|
||||
if field not in BASE_FIELDS:
|
||||
raise KeyError('Invalid column: ' + str(field))
|
||||
raise KeyError('Invalid column: {}'.format(field))
|
||||
|
||||
if isinstance(assets, collections.Iterable):
|
||||
values = list()
|
||||
for asset in assets:
|
||||
value = self.get_single_spot_value(
|
||||
asset, field, data_frequency)
|
||||
values.append(value)
|
||||
values = []
|
||||
for asset in assets:
|
||||
value = self.get_single_spot_value(asset, field, data_frequency)
|
||||
values.append(value)
|
||||
|
||||
return values
|
||||
else:
|
||||
return self.get_single_spot_value(
|
||||
assets, field, data_frequency)
|
||||
return values
|
||||
|
||||
def get_single_spot_value(self, asset, field, data_frequency):
|
||||
"""
|
||||
|
||||
@@ -107,6 +107,6 @@ class ExchangeDataPortalTestCase:
|
||||
|
||||
date = pd.to_datetime('2017-09-10', utc=True)
|
||||
value = self.data_portal_backtest.get_spot_value(
|
||||
assets[0], 'close', date, 'minute')
|
||||
assets, 'close', date, 'minute')
|
||||
pass
|
||||
|
||||
|
||||
Reference in New Issue
Block a user