diff --git a/tests/pipeline/base.py b/tests/pipeline/base.py index 06ddae44..17556087 100644 --- a/tests/pipeline/base.py +++ b/tests/pipeline/base.py @@ -93,7 +93,7 @@ class BasePipelineTestCase(TestCase): Mapping from termname -> computed result. """ engine = SimplePipelineEngine( - ExplodingObject(), + lambda column: ExplodingObject(), self.__calendar, self.__finder, ) diff --git a/tests/pipeline/test_engine.py b/tests/pipeline/test_engine.py index db89e26a..776881b2 100644 --- a/tests/pipeline/test_engine.py +++ b/tests/pipeline/test_engine.py @@ -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() diff --git a/tests/pipeline/test_pipeline_algo.py b/tests/pipeline/test_pipeline_algo.py index 40fef5f5..09ac4aed 100644 --- a/tests/pipeline/test_pipeline_algo.py +++ b/tests/pipeline/test_pipeline_algo.py @@ -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, diff --git a/zipline/algorithm.py b/zipline/algorithm.py index c6e967a3..83685601 100644 --- a/zipline/algorithm.py +++ b/zipline/algorithm.py @@ -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, ) diff --git a/zipline/pipeline/data/dataset.py b/zipline/pipeline/data/dataset.py index 211d03b0..dd64fd50 100644 --- a/zipline/pipeline/data/dataset.py +++ b/zipline/pipeline/data/dataset.py @@ -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 diff --git a/zipline/pipeline/engine.py b/zipline/pipeline/engine.py index d1d421d3..ba5485bc 100644 --- a/zipline/pipeline/engine.py +++ b/zipline/pipeline/engine.py @@ -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, )