Merge pull request #817 from quantopian/speedup-retrieve-all

Speedup retrieve all
This commit is contained in:
Scott Sanderson
2015-11-16 18:22:59 -05:00
11 changed files with 746 additions and 210 deletions
+1 -1
View File
@@ -3,7 +3,7 @@ verbosity=2
detailed-errors=1
with-ignore-docstrings=1
with-timer=1
timer-top-n=15
timer-filter=warning
[metadata]
description-file = README.rst
+2 -2
View File
@@ -12,7 +12,7 @@ from zipline.finance.trading import TradingEnvironment
from zipline.pipeline.engine import SimplePipelineEngine
from zipline.pipeline.term import AssetExists
from zipline.utils.pandas_utils import explode
from zipline.utils.test_utils import make_simple_asset_info, ExplodingObject
from zipline.utils.test_utils import make_simple_equity_info, ExplodingObject
from zipline.utils.tradingcalendar import trading_day
@@ -52,7 +52,7 @@ class BasePipelineTestCase(TestCase):
# Set up env for test
env = TradingEnvironment()
env.write_data(
equities_df=make_simple_asset_info(
equities_df=make_simple_equity_info(
assets,
self.__calendar[0],
self.__calendar[-1],
+8 -8
View File
@@ -29,18 +29,18 @@ from zipline.pipeline.loaders.blaze import (
NonPipelineField,
)
from zipline.utils.numpy_utils import repeat_last_axis
from zipline.utils.test_utils import tmp_asset_finder, make_simple_asset_info
from zipline.utils.test_utils import tmp_asset_finder, make_simple_equity_info
nameof = op.attrgetter('name')
dtypeof = op.attrgetter('dtype')
asset_infos = (
(make_simple_asset_info(
(make_simple_equity_info(
tuple(map(ord, 'ABC')),
pd.Timestamp(0),
pd.Timestamp('2015'),
),),
(make_simple_asset_info(
(make_simple_equity_info(
tuple(map(ord, 'ABCD')),
pd.Timestamp(0),
pd.Timestamp('2015'),
@@ -333,7 +333,7 @@ class BlazeToPipelineTestCase(TestCase):
dates = self.dates
asset_info = asset_infos[0][0]
with tmp_asset_finder(asset_info) as finder:
with tmp_asset_finder(equities=asset_info) as finder:
result = SimplePipelineEngine(
loader,
dates,
@@ -422,7 +422,7 @@ class BlazeToPipelineTestCase(TestCase):
expected_views,
)
with tmp_asset_finder(asset_info) as finder:
with tmp_asset_finder(equities=asset_info) as finder:
expected_output = pd.DataFrame(
list(concatv([12] * nassets, [13] * nassets, [14] * nassets)),
index=pd.MultiIndex.from_product((
@@ -466,7 +466,7 @@ class BlazeToPipelineTestCase(TestCase):
'2014-01-03': repeat_last_axis(np.array([11.0, 2.0]), nassets),
})
with tmp_asset_finder(asset_info) as finder:
with tmp_asset_finder(equities=asset_info) as finder:
expected_output = pd.DataFrame(
list(concatv([10] * nassets, [11] * nassets)),
index=pd.MultiIndex.from_product((
@@ -534,7 +534,7 @@ class BlazeToPipelineTestCase(TestCase):
pd.Timestamp('2014-01-06'),
])
with tmp_asset_finder(asset_info) as finder:
with tmp_asset_finder(equities=asset_info) as finder:
expected_output = pd.DataFrame(
expected_output_buffer,
index=pd.MultiIndex.from_product((
@@ -594,7 +594,7 @@ class BlazeToPipelineTestCase(TestCase):
# omitting the 4th and 5th to simulate a weekend
pd.Timestamp('2014-01-06'),
])
with tmp_asset_finder(asset_info) as finder:
with tmp_asset_finder(equities=asset_info) as finder:
expected_output = pd.DataFrame(
list(concatv([10] * nassets, [11] * nassets)),
index=pd.MultiIndex.from_product((
+5 -5
View File
@@ -50,8 +50,8 @@ from zipline.pipeline.factors import (
)
from zipline.utils.memoize import lazyval
from zipline.utils.test_utils import (
make_rotating_asset_info,
make_simple_asset_info,
make_rotating_equity_info,
make_simple_equity_info,
product_upper_triangle,
check_arrays,
)
@@ -151,7 +151,7 @@ class ConstantInputTestCase(TestCase):
assets=self.assets,
)
self.asset_info = make_simple_asset_info(
self.asset_info = make_simple_equity_info(
self.assets,
start_date=self.dates[0],
end_date=self.dates[-1],
@@ -498,7 +498,7 @@ class FrameInputTestCase(TestCase):
tz='UTC',
)
asset_info = make_simple_asset_info(
asset_info = make_simple_equity_info(
cls.assets,
start_date=cls.dates[0],
end_date=cls.dates[-1],
@@ -608,7 +608,7 @@ class SyntheticBcolzTestCase(TestCase):
cls.trading_day = day = cls.env.trading_day
cls.calendar = date_range('2015', '2015-08', tz='UTC', freq=day)
cls.asset_info = make_rotating_asset_info(
cls.asset_info = make_rotating_equity_info(
num_assets=6,
first_start=cls.first_asset_start,
frequency=day,
+2 -2
View File
@@ -57,7 +57,7 @@ from zipline.pipeline.loaders.equity_pricing_loader import (
USEquityPricingLoader,
)
from zipline.utils.test_utils import (
make_simple_asset_info,
make_simple_equity_info,
str_to_seconds,
)
from zipline.utils.tradingcalendar import (
@@ -332,7 +332,7 @@ class PipelineAlgorithmTestCase(TestCase):
cls.MSFT = 2
cls.BRK_A = 3
cls.assets = [cls.AAPL, cls.MSFT, cls.BRK_A]
asset_info = make_simple_asset_info(
asset_info = make_simple_equity_info(
cls.assets,
Timestamp('2014'),
Timestamp('2015'),
+214 -15
View File
@@ -28,6 +28,7 @@ import pandas as pd
from pandas.tseries.tools import normalize_date
from pandas.util.testing import assert_frame_equal
from nose_parameterized import parameterized
from numpy import full
import sqlalchemy as sa
@@ -38,6 +39,9 @@ from zipline.assets import (
AssetFinder,
AssetFinderCachedEquities,
)
from six import itervalues
from toolz import valmap
from zipline.assets.futures import (
cme_code_to_month,
FutureChain,
@@ -50,17 +54,23 @@ from zipline.assets.asset_writer import (
_version_table_schema,
)
from zipline.errors import (
SymbolNotFound,
EquitiesNotFound,
FutureContractsNotFound,
MultipleSymbolsFound,
SidAssignmentError,
RootSymbolNotFound,
AssetDBVersionError,
SidAssignmentError,
SidsNotFound,
SymbolNotFound,
)
from zipline.finance.trading import TradingEnvironment, noop_load
from zipline.utils.test_utils import (
all_subindices,
make_rotating_asset_info,
tmp_assets_db
make_commodity_future_info,
make_rotating_equity_info,
make_simple_equity_info,
tmp_assets_db,
tmp_asset_finder,
)
@@ -105,7 +115,7 @@ def build_lookup_generic_cases(asset_finder_type):
},
],
index='sid')
with tmp_assets_db(frame) as assets_db:
with tmp_assets_db(equities=frame) as assets_db:
finder = asset_finder_type(assets_db)
dupe_0, dupe_1, unique = assets = [
finder.retrieve_asset(i)
@@ -465,29 +475,72 @@ class AssetFinderTestCase(TestCase):
self.assertEqual(result.sid, i)
def test_lookup_symbol_from_multiple_valid(self):
# This test asserts that we resolve conflicts in accordance with the
# following rules when we have multiple assets holding the same symbol
# at the same time:
# If multiple SIDs exist for symbol S at time T, return the candidate
# SID whose start_date is highest. (200 cases)
# If multiple SIDs exist for symbol S at time T, the best candidate
# SIDs share the highest start_date, return the SID with the highest
# end_date. (34 cases)
# It is the opinion of the author (ssanderson) that we should consider
# this malformed input and fail here. But this is the current indended
# behavior of the code, and I accidentally broke it while refactoring.
# These will serve as regression tests until the time comes that we
# decide to enforce this as an error.
# See https://github.com/quantopian/zipline/issues/837 for more
# details.
df = pd.DataFrame.from_records(
[
{
'sid': 1,
'symbol': 'multiple',
'start_date': pd.Timestamp('2010-01-01'),
'end_date': pd.Timestamp('2013-01-01'),
'end_date': pd.Timestamp('2012-01-01'),
'exchange': 'NYSE'
},
# Same as asset 1, but with a later end date.
{
'sid': 2,
'symbol': 'multiple',
'start_date': pd.Timestamp('2012-01-01'),
'end_date': pd.Timestamp('2014-01-01'),
'start_date': pd.Timestamp('2010-01-01'),
'end_date': pd.Timestamp('2013-01-01'),
'exchange': 'NYSE'
}
},
# Same as asset 1, but with a later start_date
{
'sid': 3,
'symbol': 'multiple',
'start_date': pd.Timestamp('2011-01-01'),
'end_date': pd.Timestamp('2012-01-01'),
'exchange': 'NYSE'
},
]
)
self.env.write_data(equities_df=df)
finder = self.asset_finder_type(self.env.engine)
result = finder.lookup_symbol('MULTIPLE', pd.Timestamp('2012-05-05'))
self.assertEqual(result.symbol, 'MULTIPLE')
self.assertEqual(result.sid, 2)
def check(expected_sid, date):
result = finder.lookup_symbol(
'MULTIPLE', date,
)
self.assertEqual(result.symbol, 'MULTIPLE')
self.assertEqual(result.sid, expected_sid)
with tmp_asset_finder(finder_cls=self.asset_finder_type,
equities=df) as finder:
self.assertIsInstance(finder, self.asset_finder_type)
# Sids 1 and 2 are eligible here. We should get asset 2 because it
# has the later end_date.
check(2, pd.Timestamp('2010-12-31'))
# Sids 1, 2, and 3 are eligible here. We should get sid 3 because
# it has a later start_date
check(3, pd.Timestamp('2011-01-01'))
def test_lookup_generic(self):
"""
@@ -774,7 +827,7 @@ class AssetFinderTestCase(TestCase):
trading_day = self.env.trading_day
first_start = pd.Timestamp('2015-04-01', tz='UTC')
frame = make_rotating_asset_info(
frame = make_rotating_equity_info(
num_assets=num_assets,
first_start=first_start,
frequency=self.env.trading_day,
@@ -838,6 +891,152 @@ class AssetFinderTestCase(TestCase):
self.assertTrue(2 in sids)
self.assertTrue(3 in sids)
def test_group_by_type(self):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
futures = make_commodity_future_info(
first_sid=6,
root_symbols=['CL'],
years=[2014],
)
# Intersecting sid queries, to exercise loading of partially-cached
# results.
queries = [
([0, 1, 3], [6, 7]),
([0, 2, 3], [7, 10]),
(list(equities.index), list(futures.index)),
]
with tmp_asset_finder(equities=equities, futures=futures) as finder:
for equity_sids, future_sids in queries:
results = finder.group_by_type(equity_sids + future_sids)
self.assertEqual(
results,
{'equity': set(equity_sids), 'future': set(future_sids)},
)
@parameterized.expand([
(Equity, 'retrieve_equities', EquitiesNotFound),
(Future, 'retrieve_futures_contracts', FutureContractsNotFound),
])
def test_retrieve_specific_type(self, type_, lookup_name, failure_type):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
max_equity = equities.index.max()
futures = make_commodity_future_info(
first_sid=max_equity + 1,
root_symbols=['CL'],
years=[2014],
)
equity_sids = [0, 1]
future_sids = [max_equity + 1, max_equity + 2, max_equity + 3]
if type_ == Equity:
success_sids = equity_sids
fail_sids = future_sids
else:
fail_sids = equity_sids
success_sids = future_sids
with tmp_asset_finder(equities=equities, futures=futures) as finder:
# Run twice to exercise caching.
lookup = getattr(finder, lookup_name)
for _ in range(2):
results = lookup(success_sids)
self.assertIsInstance(results, dict)
self.assertEqual(set(results.keys()), set(success_sids))
self.assertEqual(
valmap(int, results),
dict(zip(success_sids, success_sids)),
)
self.assertEqual(
{type_},
{type(asset) for asset in itervalues(results)},
)
with self.assertRaises(failure_type):
lookup(fail_sids)
with self.assertRaises(failure_type):
# Should fail if **any** of the assets are bad.
lookup([success_sids[0], fail_sids[0]])
def test_retrieve_all(self):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
max_equity = equities.index.max()
futures = make_commodity_future_info(
first_sid=max_equity + 1,
root_symbols=['CL'],
years=[2014],
)
with tmp_asset_finder(equities=equities, futures=futures) as finder:
all_sids = finder.sids
self.assertEqual(len(all_sids), len(equities) + len(futures))
queries = [
# Empty Query.
(),
# Only Equities.
tuple(equities.index[:2]),
# Only Futures.
tuple(futures.index[:3]),
# Mixed, all cache misses.
tuple(equities.index[2:]) + tuple(futures.index[3:]),
# Mixed, all cache hits.
tuple(equities.index[2:]) + tuple(futures.index[3:]),
# Everything.
all_sids,
all_sids,
]
for sids in queries:
equity_sids = [i for i in sids if i <= max_equity]
future_sids = [i for i in sids if i > max_equity]
results = finder.retrieve_all(sids)
self.assertEqual(sids, tuple(map(int, results)))
self.assertEqual(
[Equity for _ in equity_sids] +
[Future for _ in future_sids],
list(map(type, results)),
)
self.assertEqual(
(
list(equities.symbol.loc[equity_sids]) +
list(futures.symbol.loc[future_sids])
),
list(asset.symbol for asset in results),
)
@parameterized.expand([
(EquitiesNotFound, 'equity', 'equities'),
(FutureContractsNotFound, 'future contract', 'future contracts'),
(SidsNotFound, 'asset', 'assets'),
])
def test_error_message_plurality(self,
error_type,
singular,
plural):
try:
raise error_type(sids=[1])
except error_type as e:
self.assertEqual(
str(e),
"No {singular} found for sid: 1.".format(singular=singular)
)
try:
raise error_type(sids=[1, 2])
except error_type as e:
self.assertEqual(
str(e),
"No {plural} found for sids: [1, 2].".format(plural=plural)
)
class AssetFinderCachedEquitiesTestCase(AssetFinderTestCase):
+10 -4
View File
@@ -21,8 +21,10 @@ from pandas.tseries.tools import normalize_date
import numpy as np
from datetime import datetime
from itertools import groupby, chain, repeat
from numbers import Integral
from operator import attrgetter
from six.moves import filter
from six import (
exec_,
@@ -30,7 +32,6 @@ from six import (
itervalues,
string_types,
)
from operator import attrgetter
from zipline.errors import (
@@ -608,8 +609,7 @@ class TradingAlgorithm(object):
if isinstance(identifier, Asset):
asset = self.asset_finder.retrieve_asset(sid=identifier.sid,
default_none=True)
elif hasattr(identifier, '__int__'):
elif isinstance(identifier, Integral):
asset = self.asset_finder.retrieve_asset(sid=identifier,
default_none=True)
if asset is None:
@@ -618,6 +618,12 @@ class TradingAlgorithm(object):
self.trading_environment.write_data(
equities_identifiers=identifiers_to_build)
# We need to clear out any cache misses that were stored while trying
# to do lookups. The real fix for this problem is to not construct an
# AssetFinder until we `run()` when we actually have all the data we
# need to so.
self.asset_finder._reset_caches()
return self.asset_finder.map_identifier_index_to_sids(
identifiers, as_of_date,
)
+308 -146
View File
@@ -15,22 +15,23 @@
from abc import ABCMeta
from numbers import Integral
from operator import itemgetter
import warnings
from logbook import Logger
import numpy as np
import pandas as pd
from pandas.tseries.tools import normalize_date
from pandas import isnull
from six import with_metaclass, string_types, viewkeys
from six.moves import map as imap
import sqlalchemy as sa
from toolz import compose
from zipline.errors import (
EquitiesNotFound,
FutureContractsNotFound,
MapAssetIdentifierIndexError,
MultipleSymbolsFound,
RootSymbolNotFound,
SidNotFound,
SidsNotFound,
SymbolNotFound,
MapAssetIdentifierIndexError,
)
from zipline.assets import (
Asset, Equity, Future,
@@ -41,6 +42,7 @@ from zipline.assets.asset_writer import (
ASSET_DB_VERSION,
asset_db_table_names,
)
from zipline.utils.control_flow import invert
log = Logger('assets.py')
@@ -63,13 +65,14 @@ _asset_timestamp_fields = frozenset({
})
def _convert_asset_timestamp_fields(dict):
def _convert_asset_timestamp_fields(dict_):
"""
Takes in a dict of Asset init args and converts dates to pd.Timestamps
"""
for key in (_asset_timestamp_fields & viewkeys(dict)):
value = pd.Timestamp(dict[key], tz='UTC')
dict[key] = None if pd.isnull(value) else value
for key in (_asset_timestamp_fields & viewkeys(dict_)):
value = pd.Timestamp(dict_[key], tz='UTC')
dict_[key] = None if isnull(value) else value
return dict_
class AssetFinder(object):
@@ -96,108 +99,254 @@ class AssetFinder(object):
# routing.
#
# The caches are read through, i.e. accessing an asset through
# retrieve_asset, _retrieve_equity etc. will populate the cache on
# first retrieval.
self._asset_cache = {}
self._equity_cache = {}
self._future_cache = {}
self._asset_type_cache = {}
# retrieve_asset will populate the cache on first retrieval.
self._caches = (self._asset_cache, self._asset_type_cache) = {}, {}
# Populated on first call to `lifetimes`.
self._asset_lifetimes = None
def asset_type_by_sid(self, sid):
def _reset_caches(self):
"""
Retrieve the asset type of a given sid.
Reset our asset caches.
You probably shouldn't call this method.
"""
try:
return self._asset_type_cache[sid]
except KeyError:
pass
# This method exists as a workaround for the in-place mutating behavior
# of `TradingAlgorithm._write_and_map_id_index_to_sids`. No one else
# should be calling this.
for cache in self._caches:
cache.clear()
asset_type = sa.select((self.asset_router.c.asset_type,)).where(
self.asset_router.c.sid == int(sid),
).scalar()
def lookup_asset_types(self, sids):
"""
Retrieve asset types for a list of sids.
if asset_type is not None:
self._asset_type_cache[sid] = asset_type
return asset_type
Parameters
----------
sids : list[int]
Returns
-------
types : dict[sid -> str or None]
Asset types for the provided sids.
"""
found, missing = {}, set()
for sid in sids:
try:
found[sid] = self._asset_type_cache[sid]
except KeyError:
missing.add(sid)
if not missing:
return found
router_cols = self.asset_router.c
query = sa.select((router_cols.sid, router_cols.asset_type)).where(
self.asset_router.c.sid.in_(map(int, missing))
)
for sid, type_ in query.execute().fetchall():
missing.remove(sid)
found[sid] = self._asset_type_cache[sid] = type_
for sid in missing:
found[sid] = self._asset_type_cache[sid] = None
return found
def group_by_type(self, sids):
"""
Group a list of sids by asset type.
Parameters
----------
sids : list[int]
Returns
-------
types : dict[str or None -> list[int]]
A dict mapping unique asset types to lists of sids drawn from sids.
If we fail to look up an asset, we assign it a key of None.
"""
return invert(self.lookup_asset_types(sids))
def retrieve_asset(self, sid, default_none=False):
"""
Retrieve the Asset object of a given sid.
Retrieve the Asset for a given sid.
"""
if isinstance(sid, Asset):
return sid
try:
asset = self._asset_cache[sid]
except KeyError:
asset_type = self.asset_type_by_sid(sid)
if asset_type == 'equity':
asset = self._retrieve_equity(sid)
elif asset_type == 'future':
asset = self._retrieve_futures_contract(sid)
else:
asset = None
# Cache the asset if it has been retrieved
if asset is not None:
self._asset_cache[sid] = asset
if asset is not None:
return asset
elif default_none:
return None
else:
raise SidNotFound(sid=sid)
return self.retrieve_all((sid,), default_none=default_none)[0]
def retrieve_all(self, sids, default_none=False):
return [self.retrieve_asset(sid, default_none) for sid in sids]
"""
Retrieve all assets in `sids`.
Parameters
----------
sids : interable of int
Assets to retrieve.
default_none : bool
If True, return None for failed lookups.
If False, raise `SidsNotFound`.
Returns
-------
assets : list[int or None]
A list of the same length as `sids` containing Assets (or Nones)
corresponding to the requested sids.
Raises
------
SidsNotFound
When a requested sid is not found and default_none=False.
"""
hits, missing, failures = {}, set(), []
for sid in sids:
try:
asset = self._asset_cache[sid]
if not default_none and asset is None:
# Bail early if we've already cached that we don't know
# about an asset.
raise SidsNotFound(sids=[sid])
hits[sid] = asset
except KeyError:
missing.add(sid)
# All requests were cache hits. Return requested sids in order.
if not missing:
return [hits[sid] for sid in sids]
update_hits = hits.update
# Look up cache misses by type.
type_to_assets = self.group_by_type(missing)
# Handle failures
failures = {failure: None for failure in type_to_assets.pop(None, ())}
update_hits(failures)
self._asset_cache.update(failures)
if failures and not default_none:
raise SidsNotFound(sids=list(failures))
# We don't update the asset cache here because it should already be
# updated by `self.retrieve_equities`.
update_hits(self.retrieve_equities(type_to_assets.pop('equity', ())))
update_hits(
self.retrieve_futures_contracts(type_to_assets.pop('future', ()))
)
# We shouldn't know about any other asset types.
if type_to_assets:
raise AssertionError(
"Found asset types: %s" % list(type_to_assets.keys())
)
return [hits[sid] for sid in sids]
def retrieve_equities(self, sids):
"""
Retrieve Equity objects for a list of sids.
Users generally shouldn't need to this method (instead, they should
prefer the more general/friendly `retrieve_assets`), but it has a
documented interface and tests because it's used upstream.
Parameters
----------
sids : iterable[int]
Returns
-------
equities : dict[int -> Equity]
Raises
------
EquitiesNotFound
When any requested asset isn't found.
"""
return self._retrieve_assets(sids, self.equities, Equity)
def _retrieve_equity(self, sid):
"""
Retrieve the Equity object of a given sid.
"""
return self._retrieve_asset(
sid, self._equity_cache, self.equities, Equity,
)
return self.retrieve_equities((sid,))[sid]
def _retrieve_futures_contract(self, sid):
def retrieve_futures_contracts(self, sids):
"""
Retrieve the Future object of a given sid.
Retrieve Future objects for an iterable of sids.
Users generally shouldn't need to this method (instead, they should
prefer the more general/friendly `retrieve_assets`), but it has a
documented interface and tests because it's used upstream.
Parameters
----------
sids : iterable[int]
Returns
-------
equities : dict[int -> Equity]
Raises
------
EquitiesNotFound
When any requested asset isn't found.
"""
return self._retrieve_asset(
sid, self._future_cache, self.futures_contracts, Future,
)
return self._retrieve_assets(sids, self.futures_contracts, Future)
@staticmethod
def _select_asset_by_sid(asset_tbl, sid):
return sa.select([asset_tbl]).where(asset_tbl.c.sid == int(sid))
def _select_assets_by_sid(asset_tbl, sids):
return sa.select([asset_tbl]).where(
asset_tbl.c.sid.in_(map(int, sids))
)
@staticmethod
def _select_asset_by_symbol(asset_tbl, symbol):
return sa.select([asset_tbl]).where(asset_tbl.c.symbol == symbol)
def _retrieve_asset(self, sid, cache, asset_tbl, asset_type):
try:
return cache[sid]
except KeyError:
pass
def _retrieve_assets(self, sids, asset_tbl, asset_type):
"""
Internal function for loading assets from a table.
data = self._select_asset_by_sid(asset_tbl, sid).execute().fetchone()
# Convert 'data' from a RowProxy object to a dict, to allow assignment
data = dict(data.items())
if data:
_convert_asset_timestamp_fields(data)
This should be the only method of `AssetFinder` that writes Assets into
self._asset_cache.
asset = asset_type(**data)
else:
asset = None
Parameters
---------
sids : iterable of int
Asset ids to look up.
asset_tbl : sqlalchemy.Table
Table from which to query assets.
asset_type : type
Type of asset to be constructed.
cache[sid] = asset
return asset
Returns
-------
assets : dict[int -> Asset]
Dict mapping requested sids to the retrieved assets.
"""
# Fastpath for empty request.
if not sids:
return {}
cache = self._asset_cache
hits = {}
# Load misses from the db.
query = self._select_assets_by_sid(asset_tbl, sids)
for row in imap(dict, query.execute().fetchall()):
asset = asset_type(**_convert_asset_timestamp_fields(row))
sid = asset.sid
hits[sid] = cache[sid] = asset
# If we get here, it means something in our code thought that a
# particular sid was an equity/future and called this function with a
# concrete type, but we couldn't actually resolve the asset. This is
# an error in our code, not a user-input error.
misses = tuple(set(sids) - viewkeys(hits))
if misses:
if asset_type == Equity:
raise EquitiesNotFound(sids=misses)
else:
raise FutureContractsNotFound(sids=misses)
return hits
def _get_fuzzy_candidates(self, fuzzy_symbol):
candidates = sa.select(
@@ -275,10 +424,9 @@ class AssetFinder(object):
return self._retrieve_equity(candidates[0]['sid'])
def _get_equities_from_candidates(self, candidates):
return list(map(
compose(self._retrieve_equity, itemgetter('sid')),
candidates,
))
sids = map(itemgetter('sid'), candidates)
results = self.retrieve_equities(sids)
return [results[sid] for sid in sids]
def lookup_symbol(self, symbol, as_of_date, fuzzy=False):
"""
@@ -289,12 +437,11 @@ class AssetFinder(object):
If no Equity was active at as_of_date raises SymbolNotFound.
"""
company_symbol, share_class_symbol, fuzzy_symbol = \
split_delimited_symbol(symbol)
if as_of_date:
# Format inputs
as_of_date = pd.Timestamp(normalize_date(as_of_date))
as_of_date = pd.Timestamp(as_of_date).normalize()
ad_value = as_of_date.value
if fuzzy:
@@ -379,22 +526,7 @@ class AssetFinder(object):
# If no data found, raise an exception
if not data:
raise SymbolNotFound(symbol=symbol)
# If we find a contract, check whether it's been cached
try:
return self._future_cache[data['sid']]
except KeyError:
pass
# Build the Future object from its parameters
data = dict(data.items())
_convert_asset_timestamp_fields(data)
future = Future(**data)
# Cache the Future object.
self._future_cache[data['sid']] = future
return future
return self.retrieve_asset(data['sid'])
def lookup_future_chain(self, root_symbol, as_of_date):
""" Return the futures chain for a given root symbol.
@@ -490,7 +622,8 @@ class AssetFinder(object):
if count == 0:
raise RootSymbolNotFound(root_symbol=root_symbol)
return list(map(self._retrieve_futures_contract, sids))
contracts = self.retrieve_futures_contracts(sids)
return [contracts[sid] for sid in sids]
@property
def sids(self):
@@ -516,7 +649,7 @@ class AssetFinder(object):
elif isinstance(asset_convertible, Integral):
try:
result = self.retrieve_asset(int(asset_convertible))
except SidNotFound:
except SidsNotFound:
missing.append(asset_convertible)
return None
matches.append(result)
@@ -566,7 +699,7 @@ class AssetFinder(object):
return matches[0], missing
except IndexError:
if hasattr(asset_convertible_or_iterable, '__int__'):
raise SidNotFound(sid=asset_convertible_or_iterable)
raise SidsNotFound(sids=[asset_convertible_or_iterable])
else:
raise SymbolNotFound(symbol=asset_convertible_or_iterable)
@@ -623,9 +756,8 @@ class AssetFinder(object):
self._lookup_generic_scalar(identifier, as_of_date,
matches, missing)
# Handle missing assets
if len(missing) > 0:
warnings.warn("Missing assets for identifiers: %s" % missing)
if missing:
raise ValueError("Missing assets for identifiers: %s" % missing)
# Return a list of the sids of the found assets
return [asset.sid for asset in matches]
@@ -767,57 +899,39 @@ class AssetFinderCachedEquities(AssetFinder):
fuzzy_symbol, []
).append(asset)
def _convert_row_to_equity(self, equity):
def _convert_row_to_equity(self, row):
"""
Converts a SQLAlchemy equity row to an Equity object.
"""
data = dict(equity.items())
_convert_asset_timestamp_fields(data)
asset = Equity(**data)
return asset
return Equity(**_convert_asset_timestamp_fields(dict(row)))
def _get_fuzzy_candidates(self, fuzzy_symbol):
if fuzzy_symbol in self.fuzzy_symbol_hashed_equities:
return self.fuzzy_symbol_hashed_equities[fuzzy_symbol]
return []
return self.fuzzy_symbol_hashed_equities.get(fuzzy_symbol, ())
def _get_fuzzy_candidates_in_range(self, fuzzy_symbol, ad_value):
equities = self._get_fuzzy_candidates(fuzzy_symbol)
fuzzy_candidates = []
for equity in equities:
if (equity.start_date.value <=
ad_value <=
equity.end_date.value):
fuzzy_candidates.append(equity)
return fuzzy_candidates
return only_active_assets(
ad_value,
self._get_fuzzy_candidates(fuzzy_symbol),
)
def _get_split_candidates(self, company_symbol, share_class_symbol):
if (company_symbol, share_class_symbol) in \
self.company_share_class_hashed_equities:
return self.company_share_class_hashed_equities[(
company_symbol, share_class_symbol)]
return []
return self.company_share_class_hashed_equities.get(
(company_symbol, share_class_symbol),
(),
)
def _get_split_candidates_in_range(self,
company_symbol,
share_class_symbol,
ad_value):
equities = self._get_split_candidates(
company_symbol, share_class_symbol
return sorted(
only_active_assets(
ad_value,
self._get_split_candidates(company_symbol, share_class_symbol),
),
key=lambda x: (x.start_date, x.end_date),
reverse=True,
)
best_candidates = []
for equity in equities:
if (equity.start_date.value <=
ad_value <=
equity.end_date.value):
best_candidates.append(equity)
if best_candidates:
best_candidates = sorted(
best_candidates,
key=lambda x: (x.start_date, x.end_date),
reverse=True
)
return best_candidates
def _resolve_no_matching_candidates(self,
company_symbol,
@@ -843,3 +957,51 @@ class AssetFinderCachedEquities(AssetFinder):
def _get_equities_from_candidates(self, candidates):
return candidates
def was_active(reference_date_value, asset):
"""
Whether or not `asset` was active at the time corresponding to
`reference_date_value`.
Parameters
----------
reference_date_value : int
Date, represented as nanoseconds since EPOCH, for which we want to know
if `asset` was alive. This is generally the result of accessing the
`value` attribute of a pandas Timestamp.
asset : Asset
The asset object to check.
Returns
-------
was_active : bool
Whether or not the `asset` existed at the specified time.
"""
return (
asset.start_date.value
<= reference_date_value
<= asset.end_date.value
)
def only_active_assets(reference_date_value, assets):
"""
Filter an iterable of Asset objects down to just assets that were alive at
the time corresponding to `reference_date_value`.
Parameters
----------
reference_date_value : int
Date, represented as nanoseconds since EPOCH, for which we want to know
if `asset` was alive. This is generally the result of accessing the
`value` attribute of a pandas Timestamp.
assets : iterable[Asset]
The assets to filter.
Returns
-------
active_assets : list
List of the active assets from `assets` on the requested date.
"""
return [a for a in assets if was_active(reference_date_value, a)]
+41 -7
View File
@@ -13,12 +13,13 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from zipline.utils.memoize import lazyval
class ZiplineError(Exception):
msg = None
def __init__(self, *args, **kwargs):
self.args = args
def __init__(self, **kwargs):
self.kwargs = kwargs
self.message = str(self)
@@ -231,13 +232,46 @@ Root symbol '{root_symbol}' was not found.
""".strip()
class SidNotFound(ZiplineError):
class SidsNotFound(ZiplineError):
"""
Raised when a retrieve_asset() call contains a non-existent sid.
Raised when a retrieve_asset() or retrieve_all() call contains a
non-existent sid.
"""
msg = """
Asset with sid '{sid}' was not found.
""".strip()
@lazyval
def plural(self):
return len(self.sids) > 1
@lazyval
def sids(self):
return self.kwargs['sids']
@lazyval
def msg(self):
if self.plural:
return "No assets found for sids: {sids}."
return "No asset found for sid: {sids[0]}."
class EquitiesNotFound(SidsNotFound):
"""
Raised when a call to `retrieve_equities` fails to find an asset.
"""
@lazyval
def msg(self):
if self.plural:
return "No equities found for sids: {sids}."
return "No equity found for sid: {sids[0]}."
class FutureContractsNotFound(SidsNotFound):
"""
Raised when a call to `retrieve_futures_contracts` fails to find an asset.
"""
@lazyval
def msg(self):
if self.plural:
return "No future contracts found for sids: {sids}."
return "No future contract found for sid: {sids[0]}."
class ConsumeAssetMetaDataError(ZiplineError):
+17
View File
@@ -1,6 +1,7 @@
"""
Control flow utilities.
"""
from six import iteritems
from warnings import (
catch_warnings,
filterwarnings,
@@ -54,3 +55,19 @@ def ignore_nanwarnings():
{'category': RuntimeWarning, 'module': 'numpy.lib.nanfunctions'},
)
)
def invert(d):
"""
Invert a dictionary into a dictionary of sets.
>>> invert({'a': 1, 'b': 2, 'c': 1})
{1: {'a', 'c'}, 2: {'b'}}
"""
out = {}
for k, v in iteritems(d):
try:
out[v].add(k)
except KeyError:
out[v] = {k}
return out
+138 -20
View File
@@ -14,12 +14,14 @@ from logbook import FileHandler
from mock import patch
from numpy.testing import assert_allclose, assert_array_equal
import pandas as pd
from six import itervalues
from pandas.tseries.offsets import MonthBegin
from six import iteritems, itervalues
from six.moves import filter
from sqlalchemy import create_engine
from zipline.assets import AssetFinder
from zipline.assets.asset_writer import AssetDBWriterFromDataFrame
from zipline.assets.futures import CME_CODE_TO_MONTH
from zipline.finance.blotter import ORDER_STATUS
from zipline.utils import security_list
@@ -233,11 +235,11 @@ def all_subindices(index):
)
def make_rotating_asset_info(num_assets,
first_start,
frequency,
periods_between_starts,
asset_lifetime):
def make_rotating_equity_info(num_assets,
first_start,
frequency,
periods_between_starts,
asset_lifetime):
"""
Create a DataFrame representing lifetimes of assets that are constantly
rotating in and out of existence.
@@ -262,7 +264,6 @@ def make_rotating_asset_info(num_assets,
"""
return pd.DataFrame(
{
'sid': range(num_assets),
'symbol': [chr(ord('A') + i) for i in range(num_assets)],
# Start a new asset every `periods_between_starts` days.
'start_date': pd.date_range(
@@ -277,11 +278,12 @@ def make_rotating_asset_info(num_assets,
periods=num_assets,
),
'exchange': 'TEST',
}
},
index=range(num_assets),
)
def make_simple_asset_info(assets, start_date, end_date, symbols=None):
def make_simple_equity_info(assets, start_date, end_date, symbols=None):
"""
Create a DataFrame representing assets that exist for the full duration
between `start_date` and `end_date`.
@@ -305,12 +307,122 @@ def make_simple_asset_info(assets, start_date, end_date, symbols=None):
symbols = list(ascii_uppercase[:num_assets])
return pd.DataFrame(
{
'sid': assets,
'symbol': symbols,
'start_date': [start_date] * num_assets,
'end_date': [end_date] * num_assets,
'exchange': 'TEST',
}
},
index=assets,
)
def make_future_info(first_sid,
root_symbols,
years,
notice_date_func,
expiration_date_func,
start_date_func,
month_codes=None):
"""
Create a DataFrame representing futures for `root_symbols` during `year`.
Generates a contract per triple of (symbol, year, month) supplied to
`root_symbols`, `years`, and `month_codes`.
Parameters
----------
first_sid : int
The first sid to use for assigning sids to the created contracts.
root_symbols : list[str]
A list of root symbols for which to create futures.
years : list[int or str]
Years (e.g. 2014), for which to produce individual contracts.
notice_date_func : (Timestamp) -> Timestamp
Function to generate notice dates from first of the month associated
with asset month code. Return NaT to simulate futures with no notice
date.
expiration_date_func : (Timestamp) -> Timestamp
Function to generate expiration dates from first of the month
associated with asset month code.
start_date_func : (Timestamp) -> Timestamp, optional
Function to generate start dates from first of the month associated
with each asset month code. Defaults to a start_date one year prior
to the month_code date.
month_codes : dict[str -> [1..12]], optional
Dictionary of month codes for which to create contracts. Entries
should be strings mapped to values from 1 (January) to 12 (December).
Default is zipline.futures.CME_CODE_TO_MONTH
Returns
-------
futures_info : pd.DataFrame
DataFrame of futures data suitable for passing to an
AssetDBWriterFromDataFrame.
"""
if month_codes is None:
month_codes = CME_CODE_TO_MONTH
year_strs = list(map(str, years))
years = [pd.Timestamp(s, tz='UTC') for s in year_strs]
# Pairs of string/date like ('K06', 2006-05-01)
contract_suffix_to_beginning_of_month = tuple(
(month_code + year_str[-2:], year + MonthBegin(month_num))
for ((year, year_str), (month_code, month_num))
in product(
zip(years, year_strs),
iteritems(month_codes),
)
)
contracts = []
parts = product(root_symbols, contract_suffix_to_beginning_of_month)
for sid, (root_sym, (suffix, month_begin)) in enumerate(parts, first_sid):
contracts.append({
'sid': sid,
'root_symbol': root_sym,
'symbol': root_sym + suffix,
'start_date': start_date_func(month_begin),
'notice_date': notice_date_func(month_begin),
'expiration_date': notice_date_func(month_begin),
'contract_multiplier': 500,
})
return pd.DataFrame.from_records(contracts, index='sid').convert_objects()
def make_commodity_future_info(first_sid,
root_symbols,
years,
month_codes=None):
"""
Make futures testing data that simulates the notice/expiration date
behavior of physical commodities like oil.
Parameters
----------
first_sid : int
root_symbols : list[str]
years : list[int]
month_codes : dict[str -> int]
Expiration dates are on the 20th of the month prior to the month code.
Notice dates are are on the 20th two months prior to the month code.
Start dates are one year before the contract month.
See Also
--------
make_future_info
"""
nineteen_days = pd.Timedelta(days=19)
one_year = pd.Timedelta(days=365)
return make_future_info(
first_sid=first_sid,
root_symbols=root_symbols,
years=years,
notice_date_func=lambda dt: dt - MonthBegin(2) + nineteen_days,
expiration_date_func=lambda dt: dt - MonthBegin(1) + nineteen_days,
start_date_func=lambda dt: dt - one_year,
month_codes=month_codes,
)
@@ -372,15 +484,17 @@ class tmp_assets_db(object):
The data to feed to the writer. By default this maps:
('A', 'B', 'C') -> map(ord, 'ABC')
"""
def __init__(self, data=None):
def __init__(self, **frames):
self._eng = None
self._data = AssetDBWriterFromDataFrame(
data if data is not None else make_simple_asset_info(
list(map(ord, 'ABC')),
pd.Timestamp(0),
pd.Timestamp('2015'),
)
)
if not frames:
frames = {
'equities': make_simple_equity_info(
list(map(ord, 'ABC')),
pd.Timestamp(0),
pd.Timestamp('2015'),
)
}
self._data = AssetDBWriterFromDataFrame(**frames)
def __enter__(self):
self._eng = eng = create_engine('sqlite://')
@@ -400,8 +514,12 @@ class tmp_asset_finder(tmp_assets_db):
data : dict, optional
The data to feed to the writer
"""
def __init__(self, finder_cls=AssetFinder, **frames):
self._finder_cls = finder_cls
super(tmp_asset_finder, self).__init__(**frames)
def __enter__(self):
return AssetFinder(super(tmp_asset_finder, self).__enter__())
return self._finder_cls(super(tmp_asset_finder, self).__enter__())
class SubTestFailures(AssertionError):