TEST: Parameterize over window_length.

This commit is contained in:
Scott Sanderson
2016-07-24 21:21:40 -04:00
parent be857ead0e
commit 3cc1cf078a
+9 -8
View File
@@ -26,7 +26,7 @@ from numpy.random import randn, seed as random_seed
from zipline.errors import BadPercentileBounds from zipline.errors import BadPercentileBounds
from zipline.pipeline import Filter, Factor, TermGraph from zipline.pipeline import Filter, Factor, TermGraph
from zipline.pipeline.factors import CustomFactor from zipline.pipeline.factors import CustomFactor
from zipline.testing import check_arrays from zipline.testing import check_arrays, parameter_space
from zipline.utils.numpy_utils import float64_dtype from zipline.utils.numpy_utils import float64_dtype
from .base import BasePipelineTestCase, with_default_shape from .base import BasePipelineTestCase, with_default_shape
@@ -383,13 +383,11 @@ class FilterTestCase(BasePipelineTestCase):
) )
check_arrays(results['isfinite'], isfinite(data)) check_arrays(results['isfinite'], isfinite(data))
def test_window_safe(self): @parameter_space(factor_len=[2, 3, 4])
def test_window_safe(self, factor_len):
# all true data set of (days, securities) # all true data set of (days, securities)
data = full(self.default_shape, True, dtype=bool) data = full(self.default_shape, True, dtype=bool)
# rolling window length for TestFactor
k = 8
class InputFilter(Filter): class InputFilter(Filter):
inputs = () inputs = ()
window_length = 0 window_length = 0
@@ -397,7 +395,7 @@ class FilterTestCase(BasePipelineTestCase):
class TestFactor(CustomFactor): class TestFactor(CustomFactor):
dtype = float64_dtype dtype = float64_dtype
inputs = (InputFilter(), ) inputs = (InputFilter(), )
window_length = k window_length = factor_len
def compute(self, today, assets, out, filter_): def compute(self, today, assets, out, filter_):
# sum for each column # sum for each column
@@ -412,5 +410,8 @@ class FilterTestCase(BasePipelineTestCase):
n = self.default_shape[0] n = self.default_shape[0]
# shape of output array # shape of output array
output_shape = ((n - k + 1), self.default_shape[1]) output_shape = ((n - factor_len + 1), self.default_shape[1])
check_arrays(results['windowsafe'], full(output_shape, k)) check_arrays(
results['windowsafe'],
full(output_shape, factor_len, dtype=float64)
)