mirror of
https://github.com/wassname/catalyst.git
synced 2026-07-20 12:20:29 +08:00
PERF: Using pipeline_loader_dispatch to group by loader
instead of dataset
This commit is contained in:
@@ -93,7 +93,7 @@ class BasePipelineTestCase(TestCase):
|
||||
Mapping from termname -> computed result.
|
||||
"""
|
||||
engine = SimplePipelineEngine(
|
||||
ExplodingObject(),
|
||||
lambda column: ExplodingObject(),
|
||||
self.__calendar,
|
||||
self.__finder,
|
||||
)
|
||||
|
||||
@@ -115,7 +115,8 @@ class ConstantInputTestCase(TestCase):
|
||||
|
||||
def test_bad_dates(self):
|
||||
loader = self.loader
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
|
||||
p = Pipeline()
|
||||
|
||||
@@ -129,7 +130,8 @@ class ConstantInputTestCase(TestCase):
|
||||
loader = self.loader
|
||||
finder = self.asset_finder
|
||||
assets = array(self.assets)
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
num_dates = 5
|
||||
dates = self.dates[10:10 + num_dates]
|
||||
|
||||
@@ -152,7 +154,8 @@ class ConstantInputTestCase(TestCase):
|
||||
loader = self.loader
|
||||
finder = self.asset_finder
|
||||
assets = self.assets
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
result_shape = (num_dates, num_assets) = (5, len(assets))
|
||||
dates = self.dates[10:10 + num_dates]
|
||||
|
||||
@@ -185,7 +188,8 @@ class ConstantInputTestCase(TestCase):
|
||||
loader = self.loader
|
||||
finder = self.asset_finder
|
||||
assets = self.assets
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
shape = num_dates, num_assets = (5, len(assets))
|
||||
dates = self.dates[10:10 + num_dates]
|
||||
|
||||
@@ -228,7 +232,8 @@ class ConstantInputTestCase(TestCase):
|
||||
def test_numeric_factor(self):
|
||||
constants = self.constants
|
||||
loader = self.loader
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
num_dates = 5
|
||||
dates = self.dates[10:10 + num_dates]
|
||||
high, low = USEquityPricing.high, USEquityPricing.low
|
||||
@@ -355,7 +360,8 @@ class FrameInputTestCase(TestCase):
|
||||
high_loader = DataFrameLoader(high, high_base, adjustments)
|
||||
loader = MultiColumnLoader({low: low_loader, high: high_loader})
|
||||
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
|
||||
for window_length in range(1, 4):
|
||||
low_mavg = SimpleMovingAverage(
|
||||
@@ -465,7 +471,7 @@ class SyntheticBcolzTestCase(TestCase):
|
||||
|
||||
def test_SMA(self):
|
||||
engine = SimplePipelineEngine(
|
||||
self.pipeline_loader,
|
||||
lambda column: self.pipeline_loader,
|
||||
self.env.trading_days,
|
||||
self.finder,
|
||||
)
|
||||
@@ -517,7 +523,7 @@ class SyntheticBcolzTestCase(TestCase):
|
||||
# or zero, but verifying we correctly handle those corner cases is
|
||||
# valuable.
|
||||
engine = SimplePipelineEngine(
|
||||
self.pipeline_loader,
|
||||
lambda column: self.pipeline_loader,
|
||||
self.env.trading_days,
|
||||
self.finder,
|
||||
)
|
||||
@@ -584,7 +590,8 @@ class MultiColumnLoaderTestCase(TestCase):
|
||||
dates=self.dates,
|
||||
assets=self.assets,
|
||||
)
|
||||
engine = SimplePipelineEngine(loader, self.dates, self.asset_finder)
|
||||
engine = SimplePipelineEngine(lambda column: loader,
|
||||
self.dates, self.asset_finder)
|
||||
|
||||
sumdiff = RollingSumDifference()
|
||||
|
||||
|
||||
@@ -182,7 +182,7 @@ class ClosesOnly(TestCase):
|
||||
initialize=initialize,
|
||||
handle_data=late_attach,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.first_asset_start - trading_day,
|
||||
end=self.last_asset_end + trading_day,
|
||||
env=self.env,
|
||||
@@ -199,7 +199,7 @@ class ClosesOnly(TestCase):
|
||||
before_trading_start=late_attach,
|
||||
handle_data=barf,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.first_asset_start - trading_day,
|
||||
end=self.last_asset_end + trading_day,
|
||||
env=self.env,
|
||||
@@ -228,7 +228,7 @@ class ClosesOnly(TestCase):
|
||||
handle_data=handle_data,
|
||||
before_trading_start=before_trading_start,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.first_asset_start - trading_day,
|
||||
end=self.last_asset_end + trading_day,
|
||||
env=self.env,
|
||||
@@ -256,7 +256,7 @@ class ClosesOnly(TestCase):
|
||||
handle_data=handle_data,
|
||||
before_trading_start=before_trading_start,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.first_asset_start - trading_day,
|
||||
end=self.last_asset_end + trading_day,
|
||||
env=self.env,
|
||||
@@ -294,7 +294,7 @@ class ClosesOnly(TestCase):
|
||||
handle_data=handle_data,
|
||||
before_trading_start=before_trading_start,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.first_asset_start - trading_day,
|
||||
end=self.last_asset_end + trading_day,
|
||||
env=self.env,
|
||||
@@ -524,7 +524,7 @@ class PipelineAlgorithmTestCase(TestCase):
|
||||
handle_data=handle_data,
|
||||
before_trading_start=before_trading_start,
|
||||
data_frequency='daily',
|
||||
pipeline_loader=self.pipeline_loader,
|
||||
pipeline_loader_dispatch=lambda column: self.pipeline_loader,
|
||||
start=self.dates[max(window_lengths)],
|
||||
end=self.dates[-1],
|
||||
env=self.env,
|
||||
|
||||
@@ -232,7 +232,7 @@ class TradingAlgorithm(object):
|
||||
self.asset_finder = self.trading_environment.asset_finder
|
||||
|
||||
# Initialize Pipeline API data.
|
||||
self.init_engine(kwargs.pop('pipeline_loader', None))
|
||||
self.init_engine(kwargs.pop('pipeline_loader_dispatch', None))
|
||||
self._pipelines = {}
|
||||
# Create an always-expired cache so that we compute the first time data
|
||||
# is requested.
|
||||
@@ -323,15 +323,15 @@ class TradingAlgorithm(object):
|
||||
self.initialize_args = args
|
||||
self.initialize_kwargs = kwargs
|
||||
|
||||
def init_engine(self, loader):
|
||||
def init_engine(self, loader_dispatch):
|
||||
"""
|
||||
Construct and store a PipelineEngine from loader.
|
||||
|
||||
If loader is None, constructs a NoOpPipelineEngine.
|
||||
"""
|
||||
if loader is not None:
|
||||
if loader_dispatch is not None:
|
||||
self.engine = SimplePipelineEngine(
|
||||
loader,
|
||||
loader_dispatch,
|
||||
self.trading_environment.trading_days,
|
||||
self.asset_finder,
|
||||
)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""
|
||||
dataset.py
|
||||
"""
|
||||
from functools import total_ordering
|
||||
from six import (
|
||||
iteritems,
|
||||
with_metaclass,
|
||||
@@ -84,6 +85,7 @@ class BoundColumn(AtomicTerm):
|
||||
return self.qualname
|
||||
|
||||
|
||||
@total_ordering
|
||||
class DataSetMeta(type):
|
||||
"""
|
||||
Metaclass for DataSets
|
||||
@@ -107,6 +109,9 @@ class DataSetMeta(type):
|
||||
def columns(self):
|
||||
return self._columns
|
||||
|
||||
def __lt__(self, other):
|
||||
return id(self) < id(other)
|
||||
|
||||
|
||||
class DataSet(with_metaclass(DataSetMeta)):
|
||||
domain = None
|
||||
|
||||
@@ -92,15 +92,15 @@ class SimplePipelineEngine(object):
|
||||
which assets are in the top-level universe at any point in time.
|
||||
"""
|
||||
__slots__ = [
|
||||
'_loader',
|
||||
'_loader_dispatch',
|
||||
'_calendar',
|
||||
'_finder',
|
||||
'_root_mask_term',
|
||||
'__weakref__',
|
||||
]
|
||||
|
||||
def __init__(self, loader, calendar, asset_finder):
|
||||
self._loader = loader
|
||||
def __init__(self, loader_dispatch, calendar, asset_finder):
|
||||
self._loader_dispatch = loader_dispatch
|
||||
self._calendar = calendar
|
||||
self._finder = asset_finder
|
||||
self._root_mask_term = AssetExists()
|
||||
@@ -281,12 +281,21 @@ class SimplePipelineEngine(object):
|
||||
out.append(input_data)
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def _atomic_dataset_terms(graph, match):
|
||||
def _atomic_terms_for_loader(self, graph, loader):
|
||||
loader_dispatch = self.loader_dispatch
|
||||
for term in graph.atomic_terms:
|
||||
if term.dataset == match.dataset:
|
||||
if loader_dispatch(term) == loader:
|
||||
yield term
|
||||
|
||||
def loader_dispatch(self, term):
|
||||
if term is AssetExists():
|
||||
return None
|
||||
|
||||
loader = self._loader_dispatch(term)
|
||||
if loader is None:
|
||||
raise ValueError("Couldn't find loader for %s" % term)
|
||||
return loader
|
||||
|
||||
def compute_chunk(self, graph, dates, assets, initial_workspace):
|
||||
"""
|
||||
Compute the Pipeline terms in the graph for the requested start and end
|
||||
@@ -311,7 +320,7 @@ class SimplePipelineEngine(object):
|
||||
Dictionary mapping requested results to outputs.
|
||||
"""
|
||||
self._validate_compute_chunk_params(dates, assets, initial_workspace)
|
||||
loader = self._loader
|
||||
loader_dispatch = self.loader_dispatch
|
||||
|
||||
# Copy the supplied initial workspace so we don't mutate it in place.
|
||||
workspace = initial_workspace.copy()
|
||||
@@ -325,7 +334,9 @@ class SimplePipelineEngine(object):
|
||||
continue
|
||||
|
||||
if term.atomic:
|
||||
to_load = list(self._atomic_dataset_terms(graph, term))
|
||||
loader = loader_dispatch(term)
|
||||
to_load = sorted(self._atomic_terms_for_loader(graph, loader),
|
||||
key=lambda t: t.dataset)
|
||||
mask, mask_dates = self._mask_and_dates_for_atomic_terms(
|
||||
to_load, workspace, graph, dates,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user