DEV: Add support for specifying missing_value.

Consequently, enable support for `int`-dtyped Factors and BoundColumns.
This commit is contained in:
Scott Sanderson
2016-02-12 21:23:47 -05:00
parent 94c02c710b
commit c105735574
15 changed files with 287 additions and 104 deletions
+51 -6
View File
@@ -10,6 +10,7 @@ from numpy import (
arange,
array,
full,
where,
)
from numpy.testing import assert_array_equal
from six.moves import zip_longest
@@ -23,9 +24,21 @@ from zipline.lib.adjustment import (
from zipline.lib.adjusted_array import AdjustedArray, NOMASK
from zipline.utils.numpy_utils import (
datetime64ns_dtype,
default_missing_value_for_dtype,
float64_dtype,
int64_dtype,
make_datetime64ns,
)
from zipline.utils.test_utils import check_arrays, parameter_space
def moving_window(array, nrows):
"""
Simple moving window generator over a 2D numpy array.
"""
count = num_windows_of_length_M_on_buffers_of_length_N(nrows, len(array))
for i in range(count):
yield array[i:i + nrows]
def num_windows_of_length_M_on_buffers_of_length_N(M, N):
@@ -66,6 +79,7 @@ def _gen_unadjusted_cases(dtype):
nrows = 6
ncols = 3
data = arange(nrows * ncols).astype(dtype).reshape(nrows, ncols)
missing_value = default_missing_value_for_dtype(dtype)
for windowlen in valid_window_lengths(nrows):
@@ -78,6 +92,7 @@ def _gen_unadjusted_cases(dtype):
data,
windowlen,
{},
missing_value,
[
data[offset:offset + windowlen]
for offset in range(num_legal_windows)
@@ -230,6 +245,7 @@ def _gen_overwrite_adjustment_cases(dtype):
def _gen_expectations(baseline, adjustments, buffer_as_of, nrows):
missing_value = default_missing_value_for_dtype(baseline.dtype)
for windowlen in valid_window_lengths(nrows):
num_legal_windows = num_windows_of_length_M_on_buffers_of_length_N(
@@ -241,6 +257,7 @@ def _gen_expectations(baseline, adjustments, buffer_as_of, nrows):
baseline,
windowlen,
adjustments,
missing_value,
[
# This is a nasty expression...
#
@@ -267,9 +284,10 @@ class AdjustedArrayTestCase(TestCase):
data,
lookback,
adjustments,
missing_value,
expected):
array = AdjustedArray(data, NOMASK, adjustments)
array = AdjustedArray(data, NOMASK, adjustments, missing_value)
for _ in range(2): # Iterate 2x ensure adjusted_arrays are re-usable.
window_iter = array.traverse(lookback)
for yielded, expected_yield in zip_longest(window_iter, expected):
@@ -282,9 +300,10 @@ class AdjustedArrayTestCase(TestCase):
data,
lookback,
adjustments,
missing_value,
expected):
array = AdjustedArray(data, NOMASK, adjustments)
array = AdjustedArray(data, NOMASK, adjustments, missing_value)
for _ in range(2): # Iterate 2x ensure adjusted_arrays are re-usable.
window_iter = array.traverse(lookback)
for yielded, expected_yield in zip_longest(window_iter, expected):
@@ -301,18 +320,43 @@ class AdjustedArrayTestCase(TestCase):
data,
lookback,
adjustments,
missing_value,
expected):
array = AdjustedArray(data, NOMASK, adjustments)
array = AdjustedArray(data, NOMASK, adjustments, missing_value)
for _ in range(2): # Iterate 2x ensure adjusted_arrays are re-usable.
window_iter = array.traverse(lookback)
for yielded, expected_yield in zip_longest(window_iter, expected):
self.assertEqual(yielded.dtype, data.dtype)
assert_array_equal(yielded, expected_yield)
@parameter_space(
dtype=[float64_dtype, int64_dtype, datetime64ns_dtype],
missing_value=[0, 10000],
window_length=[2, 3],
)
def test_masking(self, dtype, missing_value, window_length):
missing_value = value_with_dtype(dtype, missing_value)
baseline_ints = arange(15).reshape(5, 3)
baseline = baseline_ints.astype(dtype)
mask = (baseline_ints % 2).astype(bool)
masked_baseline = where(mask, baseline, missing_value)
array = AdjustedArray(
baseline,
mask,
adjustments={},
missing_value=missing_value,
)
gen_expected = moving_window(masked_baseline, window_length)
gen_actual = array.traverse(window_length)
for expected, actual in zip(gen_expected, gen_actual):
check_arrays(expected, actual)
def test_invalid_lookback(self):
data = arange(30, dtype=float).reshape(6, 5)
adj_array = AdjustedArray(data, NOMASK, {})
adj_array = AdjustedArray(data, NOMASK, {}, float('nan'))
with self.assertRaises(WindowLengthTooLong):
adj_array.traverse(7)
@@ -326,7 +370,7 @@ class AdjustedArrayTestCase(TestCase):
def test_array_views_arent_writable(self):
data = arange(30, dtype=float).reshape(6, 5)
adj_array = AdjustedArray(data, NOMASK, {})
adj_array = AdjustedArray(data, NOMASK, {}, float('nan'))
for frame in adj_array.traverse(3):
with self.assertRaises(ValueError):
@@ -338,7 +382,7 @@ class AdjustedArrayTestCase(TestCase):
bad_mask = array([[0, 1, 1], [0, 0, 1]], dtype=bool)
with self.assertRaisesRegexp(ValueError, msg):
AdjustedArray(data, bad_mask, {})
AdjustedArray(data, bad_mask, {}, missing_value=-1)
def test_inspect(self):
data = arange(15, dtype=float).reshape(5, 3)
@@ -346,6 +390,7 @@ class AdjustedArrayTestCase(TestCase):
data,
NOMASK,
{4: [Float64Multiply(2, 3, 0, 0, 4.0)]},
float('nan'),
)
expected = dedent(
+43 -11
View File
@@ -31,7 +31,11 @@ from zipline.pipeline.loaders.blaze.core import (
NonPipelineField,
no_deltas_rules,
)
from zipline.utils.numpy_utils import repeat_last_axis
from zipline.utils.numpy_utils import (
float64_dtype,
int64_dtype,
repeat_last_axis,
)
from zipline.utils.test_utils import tmp_asset_finder, make_simple_equity_info
@@ -73,7 +77,8 @@ class BlazeToPipelineTestCase(TestCase):
cls.sids = sids = ord('A'), ord('B'), ord('C')
cls.df = df = pd.DataFrame({
'sid': sids * 3,
'value': (0, 1, 2, 1, 2, 3, 2, 3, 4),
'value': (0., 1., 2., 1., 2., 3., 2., 3., 4.),
'int_value': (0, 1, 2, 1, 2, 3, 2, 3, 4),
'asof_date': dates,
'timestamp': dates,
})
@@ -81,6 +86,7 @@ class BlazeToPipelineTestCase(TestCase):
var * {
sid: ?int64,
value: ?float64,
int_value: ?int64,
asof_date: datetime,
timestamp: datetime
}
@@ -91,6 +97,7 @@ class BlazeToPipelineTestCase(TestCase):
cls.macro_dshape = var * Record(dshape_)
cls.garbage_loader = BlazeLoader()
cls.missing_values = {'int_value': 0}
def test_tabular(self):
name = 'expr'
@@ -99,15 +106,20 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
)
self.assertEqual(ds.__name__, name)
self.assertTrue(issubclass(ds, DataSet))
self.assertEqual(
{c.name: c.dtype for c in ds.columns},
{'sid': np.int64, 'value': np.float64},
)
for field in ('timestamp', 'asof_date'):
self.assertIs(ds.value.dtype, float64_dtype)
self.assertIs(ds.int_value.dtype, int64_dtype)
self.assertTrue(np.isnan(ds.value.missing_value))
self.assertEqual(ds.int_value.missing_value, 0)
invalid_type_fields = ('asof_date',)
for field in invalid_type_fields:
with self.assertRaises(AttributeError) as e:
getattr(ds, field)
self.assertIn("'%s'" % field, str(e.exception))
@@ -119,6 +131,7 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
),
ds,
)
@@ -130,10 +143,11 @@ class BlazeToPipelineTestCase(TestCase):
expr.value,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
)
self.assertEqual(value.name, 'value')
self.assertIsInstance(value, BoundColumn)
self.assertEqual(value.dtype, np.float64)
self.assertIs(value.dtype, float64_dtype)
# test memoization
self.assertIs(
@@ -141,6 +155,7 @@ class BlazeToPipelineTestCase(TestCase):
expr.value,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
),
value,
)
@@ -149,6 +164,7 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
).value,
value,
)
@@ -159,6 +175,7 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=self.garbage_loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
),
value.dataset,
)
@@ -195,7 +212,11 @@ class BlazeToPipelineTestCase(TestCase):
)),
)
loader = BlazeLoader()
ds = from_blaze(expr.ds, loader=loader)
ds = from_blaze(
expr.ds,
loader=loader,
missing_values=self.missing_values,
)
self.assertEqual(len(loader), 1)
exprdata = loader[ds]
self.assertTrue(exprdata.expr.isidentical(expr.ds))
@@ -210,6 +231,7 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=loader,
no_deltas_rule=no_deltas_rules.warn,
missing_values=self.missing_values,
)
self.assertEqual(len(ws), 1)
w = ws[0].message
@@ -281,6 +303,7 @@ class BlazeToPipelineTestCase(TestCase):
expr_with_add,
deltas=None,
loader=self.garbage_loader,
missing_values=self.missing_values,
)
with self.assertRaises(TypeError):
@@ -288,6 +311,7 @@ class BlazeToPipelineTestCase(TestCase):
expr.value + 1, # put an Add in the column
deltas=None,
loader=self.garbage_loader,
missing_values=self.missing_values,
)
deltas = bz.Data(
@@ -299,6 +323,7 @@ class BlazeToPipelineTestCase(TestCase):
expr_with_add,
deltas=deltas,
loader=self.garbage_loader,
missing_values=self.missing_values,
)
with self.assertRaises(TypeError):
@@ -306,6 +331,7 @@ class BlazeToPipelineTestCase(TestCase):
expr.value + 1,
deltas=deltas,
loader=self.garbage_loader,
missing_values=self.missing_values,
)
def _test_id(self, df, dshape, expected, finder, add):
@@ -315,6 +341,7 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
)
p = Pipeline()
for a in add:
@@ -347,9 +374,11 @@ class BlazeToPipelineTestCase(TestCase):
expr,
loader=loader,
no_deltas_rule=no_deltas_rules.ignore,
missing_values=self.missing_values,
)
p = Pipeline()
p.add(ds.value.latest, 'value')
p.add(ds.int_value.latest, 'int_value')
dates = self.dates
with tmp_asset_finder() as finder:
@@ -405,7 +434,9 @@ class BlazeToPipelineTestCase(TestCase):
expected.index.levels[0],
finder.retrieve_all(expected.index.levels[1]),
))
self._test_id(self.df, self.dshape, expected, finder, ('value',))
self._test_id(
self.df, self.dshape, expected, finder, ('int_value', 'value',)
)
def test_id_ffill_out_of_window(self):
"""
@@ -512,7 +543,7 @@ class BlazeToPipelineTestCase(TestCase):
var * Record(fields),
expected,
finder,
('value', 'other'),
('value', 'int_value', 'other'),
)
def test_id_macro_dataset(self):
@@ -782,6 +813,7 @@ class BlazeToPipelineTestCase(TestCase):
deltas,
loader=loader,
no_deltas_rule=no_deltas_rules.raise_,
missing_values=self.missing_values,
)
p = Pipeline()
+6 -9
View File
@@ -42,13 +42,9 @@ class LatestTestCase(TestCase):
)
def test_latest(self):
columns = TDS.columns
pipe = Pipeline(
columns={
name: getattr(TDS, name + '_col').latest
# Intentionally not including int and bool because they're not
# yet supported.
for name in ('float', 'datetime')
}
columns={c.name: c.latest for c in columns},
)
cal_slice = slice(20, 40)
@@ -58,6 +54,7 @@ class LatestTestCase(TestCase):
dates_to_test[0],
dates_to_test[-1],
)
float_result = result.float.unstack()
expected_float_result = self.expected_latest(TDS.float_col, cal_slice)
assert_frame_equal(float_result, expected_float_result)
for column in columns:
float_result = result[column.name].unstack()
expected_float_result = self.expected_latest(column, cal_slice)
assert_frame_equal(float_result, expected_float_result)
+5 -5
View File
@@ -162,7 +162,7 @@ class ConstantInputTestCase(TestCase):
self.loader = PrecomputedLoader(
constants=self.constants,
dates=self.dates,
assets=self.asset_ids,
sids=self.asset_ids,
)
self.asset_info = make_simple_equity_info(
@@ -367,7 +367,7 @@ class ConstantInputTestCase(TestCase):
loader = PrecomputedLoader(
constants=constants,
dates=self.dates,
assets=self.asset_ids,
sids=self.asset_ids,
)
engine = SimplePipelineEngine(
lambda column: loader, self.dates, self.asset_finder,
@@ -415,7 +415,7 @@ class ConstantInputTestCase(TestCase):
def test_loader_given_multiple_columns(self):
class Loader1DataSet1(DataSet):
col1 = Column(float32)
col1 = Column(float)
col2 = Column(float32)
class Loader1DataSet2(DataSet):
@@ -433,12 +433,12 @@ class ConstantInputTestCase(TestCase):
loader1 = RecordingPrecomputedLoader(constants=constants1,
dates=self.dates,
assets=self.assets)
sids=self.assets)
constants2 = {Loader2DataSet.col1: 5,
Loader2DataSet.col2: 6}
loader2 = RecordingPrecomputedLoader(constants=constants2,
dates=self.dates,
assets=self.assets)
sids=self.assets)
engine = SimplePipelineEngine(
lambda column:
+19 -6
View File
@@ -11,7 +11,6 @@ from zipline.errors import (
InvalidDType,
TermInputsNotSpecified,
WindowLengthNotSpecified,
UnsupportedDataType,
)
from zipline.pipeline import Factor, Filter, TermGraph
from zipline.pipeline.data import Column, DataSet
@@ -21,6 +20,8 @@ from zipline.pipeline.expression import NUMEXPR_MATH_FUNCS
from zipline.utils.numpy_utils import (
datetime64ns_dtype,
float64_dtype,
int64_dtype,
NoDefaultMissingValue,
)
@@ -334,15 +335,27 @@ class ObjectIdentityTestCase(TestCase):
SomeFactor(dtype=1)
def test_latest_on_different_dtypes(self):
self.assertIsInstance(TestingDataSet.bool_col.latest, Filter)
self.assertIsInstance(TestingDataSet.float_col.latest, Factor)
self.assertIsInstance(TestingDataSet.datetime_col.latest, Factor)
self.assertIsInstance(TestingDataSet.int_col.latest, Factor)
# TODO: Support this by allowing users to provide a missing value on
# columns.
with self.assertRaises(UnsupportedDataType):
self.assertIsInstance(TestingDataSet.int_col.latest, Factor)
def test_failure_timing_on_bad_missing_values(self):
# Just constructing a bad column shouldn't fail.
Column(dtype=int64_dtype)
with self.assertRaises(NoDefaultMissingValue) as e:
class BadDataSet(DataSet):
bad_column = Column(dtype=int64_dtype)
float_column = Column(dtype=float64_dtype)
int_column = Column(dtype=int64_dtype, missing_value=3)
self.assertTrue(
str(e.exception.message).startswith(
"Failed to create Column with name 'bad_column'"
)
)
class SubDataSetTestCase(TestCase):
+28 -25
View File
@@ -17,26 +17,24 @@ from zipline.errors import (
)
from zipline.utils.numpy_utils import (
datetime64ns_dtype,
default_fillvalue_for_dtype,
float64_dtype,
int64_dtype,
uint8_dtype,
)
from zipline.utils.memoize import lazyval
from zipline.utils.sentinel import sentinel
# These class names are all the same because of our bootleg templating system.
from ._float64window import AdjustedArrayWindow as Float64Window
from ._int64window import AdjustedArrayWindow as Int64Window
from ._uint8window import AdjustedArrayWindow as UInt8Window
Infer = sentinel(
'Infer',
"Sentinel used to say 'infer missing_value from data type.'"
)
NOMASK = None
SUPPORTED_NUMERIC_DTYPES = frozenset(
map(dtype, [float32, float64, int32, int64, uint32])
FLOAT_DTYPES = frozenset(
map(dtype, [float32, float64, int32]),
)
INT_DTYPES = frozenset(
# NOTE: uint64 not supported because it can't be safely cast to int64.
map(dtype, [int32, int64, uint32]),
)
CONCRETE_WINDOW_TYPES = {
float64_dtype: Float64Window,
@@ -51,13 +49,10 @@ def _normalize_array(data):
representation, returning the coerced array and a numpy dtype object to use
as a view type when providing public view into the data.
Semantically numerical data (float*, int*, uint*) is coerced to float64 and
viewed as float64. We coerce integral data to float so that we can use NaN
as a missing value.
datetime[*] data is coerced to int64 with a viewtype of ``datetime64[ns]``.
``bool_`` data is coerced to uint8 with a viewtype of ``bool_``
- float* data is coerced to float64 with viewtype float64.
- int32, int64, and uint32 are converted to int64 with viewtype int64.
- datetime[*] data is coerced to int64 with a viewtype of datetime64[ns].
- bool_ data is coerced to uint8 with a viewtype of bool_.
Parameters
----------
@@ -70,8 +65,10 @@ def _normalize_array(data):
data_dtype = data.dtype
if data_dtype == bool_:
return data.astype(uint8), dtype(bool_)
elif data_dtype in SUPPORTED_NUMERIC_DTYPES:
elif data_dtype in FLOAT_DTYPES:
return data.astype(float64), dtype(float64)
elif data_dtype in INT_DTYPES:
return data.astype(int64), dtype(int64)
elif data_dtype.name.startswith('datetime'):
try:
outarray = data.astype('datetime64[ns]').view('int64')
@@ -105,18 +102,24 @@ class AdjustedArray(object):
adjustments : dict[int -> list[Adjustment]]
A dict mapping row indices to lists of adjustments to apply when we
reach that row.
fillvalue : object, optional
missing_value : object
A value to use to fill missing data in yielded windows.
Default behavior is to infer a value based on the dtype of `data`.
`NaN` is used for numeric data, and `NaT` is used for datetime data.
"""
__slots__ = ('_data', '_viewtype', 'adjustments', '__weakref__')
Should be a value coercible to `data.dtype`.
def __init__(self, data, mask, adjustments, fillvalue=Infer):
"""
__slots__ = (
'_data',
'_viewtype',
'adjustments',
'missing_value',
'__weakref__',
)
def __init__(self, data, mask, adjustments, missing_value):
self._data, self._viewtype = _normalize_array(data)
self.adjustments = adjustments
if fillvalue is Infer:
fillvalue = default_fillvalue_for_dtype(self.data.dtype)
self.missing_value = missing_value
if mask is not NOMASK:
if mask.dtype != bool_:
@@ -126,7 +129,7 @@ class AdjustedArray(object):
"Mask shape %s != data shape %s." %
(mask.shape, data.shape),
)
self._data[~mask] = fillvalue
self._data[~mask] = self.missing_value
@lazyval
def data(self):
+41 -7
View File
@@ -7,9 +7,13 @@ from six import (
with_metaclass,
)
from zipline.pipeline.term import Term, AssetExists
from zipline.pipeline.term import Term, AssetExists, NotSpecified
from zipline.utils.input_validation import ensure_dtype
from zipline.utils.numpy_utils import bool_dtype
from zipline.utils.numpy_utils import (
bool_dtype,
default_missing_value_for_dtype,
NoDefaultMissingValue,
)
from zipline.utils.preprocess import preprocess
@@ -19,14 +23,19 @@ class Column(object):
"""
@preprocess(dtype=ensure_dtype)
def __init__(self, dtype):
def __init__(self, dtype, missing_value=NotSpecified):
self.dtype = dtype
self.missing_value = missing_value
def bind(self, name):
"""
Bind a `Column` object to its name.
"""
return _BoundColumnDescr(dtype=self.dtype, name=name)
return _BoundColumnDescr(
dtype=self.dtype,
missing_value=self.missing_value,
name=name,
)
class _BoundColumnDescr(object):
@@ -37,8 +46,27 @@ class _BoundColumnDescr(object):
This exists so that subclasses of DataSets don't share columns with their
parent classes.
"""
def __init__(self, dtype, name):
def __init__(self, dtype, missing_value, name):
self.dtype = dtype
# Calculating missing values here guarantees that we fail quickly if
# the user fails to provide a missing value for a dtype that requires
# one (e.g. int64), but still enables us to provide an error message
# that points to the name of the failing column.
if missing_value is NotSpecified:
try:
missing_value = default_missing_value_for_dtype(dtype)
except NoDefaultMissingValue:
# Re-raise with a better message.
raise NoDefaultMissingValue(
"Failed to create Column with name {name!r} and"
" dtype {dtype} because no missing_value was provided\n\n"
"Columns with dtype {dtype} require a missing_value.\n"
"Please pass missing_value to Column() or use a different"
" dtype.".format(dtype=dtype, name=name)
)
self.missing_value = missing_value
self.name = name
def __get__(self, instance, owner):
@@ -50,6 +78,7 @@ class _BoundColumnDescr(object):
"""
return BoundColumn(
dtype=self.dtype,
missing_value=self.missing_value,
dataset=owner,
name=self.name,
)
@@ -63,11 +92,12 @@ class BoundColumn(Term):
extra_input_rows = 0
inputs = ()
def __new__(cls, dtype, dataset, name):
def __new__(cls, dtype, missing_value, dataset, name):
return super(BoundColumn, cls).__new__(
cls,
domain=dataset.domain,
dtype=dtype,
missing_value=missing_value,
dataset=dataset,
name=name,
)
@@ -106,7 +136,11 @@ class BoundColumn(Term):
from zipline.pipeline.filters import Latest
else:
from zipline.pipeline.factors import Latest
return Latest(inputs=(self,), dtype=self.dtype)
return Latest(
inputs=(self,),
dtype=self.dtype,
missing_value=self.missing_value,
)
def __repr__(self):
return "{qualname}::{dtype}".format(
+5 -2
View File
@@ -14,8 +14,11 @@ from zipline.utils.numpy_utils import (
class TestingDataSet(DataSet):
# Tell nose this isn't a test case.
__test__ = False
bool_col = Column(dtype=bool_dtype)
bool_col = Column(dtype=bool_dtype, missing_value=False)
bool_col_default_True = Column(dtype=bool_dtype, missing_value=True)
float_col = Column(dtype=float64_dtype)
datetime_col = Column(dtype=datetime64ns_dtype)
int_col = Column(dtype=int64_dtype)
int_col = Column(dtype=int64_dtype, missing_value=0)
+2 -1
View File
@@ -38,6 +38,7 @@ from zipline.utils.numpy_utils import (
bool_dtype,
datetime64ns_dtype,
float64_dtype,
int64_dtype,
)
from zipline.utils.preprocess import preprocess
@@ -303,7 +304,7 @@ def function_application(func):
return mathfunc
FACTOR_DTYPES = frozenset([datetime64ns_dtype, float64_dtype])
FACTOR_DTYPES = frozenset([datetime64ns_dtype, float64_dtype, int64_dtype])
class Factor(CompositeTerm):
+32 -7
View File
@@ -165,6 +165,7 @@ from zipline.pipeline.loaders.utils import (
normalize_data_query_bounds,
normalize_timestamp_to_query_time,
)
from zipline.pipeline.term import NotSpecified
from zipline.lib.adjusted_array import AdjustedArray
from zipline.lib.adjustment import Float64Overwrite
from zipline.utils.enum import enum
@@ -275,15 +276,21 @@ _new_names = ('BlazeDataSet_%d' % n for n in count())
@memoize
def new_dataset(expr, deltas):
"""Creates or returns a dataset from a pair of blaze expressions.
def new_dataset(expr, deltas, missing_values):
"""
Creates or returns a dataset from a pair of blaze expressions.
Parameters
----------
expr : Expr
The blaze expression representing the first known values.
The blaze expression representing the first known values.
deltas : Expr
The blaze expression representing the deltas to the data.
The blaze expression representing the deltas to the data.
missing_values : frozenset((name, value) pairs
Association pairs column name and missing_value for that column.
This needs to be a frozenset rather than a dict or tuple of tuples
because we want a collection that's unordered but still hashable.
Returns
-------
@@ -295,9 +302,16 @@ def new_dataset(expr, deltas):
This function is memoized. repeated calls with the same inputs will return
the same type.
"""
missing_values = dict(missing_values)
columns = {}
for name, type_ in expr.dshape.measure.fields:
# Don't generate a column for sid or timestamp, since they're
# implicitly the labels if the arrays that will be passed to pipeline
# Terms.
if name in (SID_FIELD_NAME, TS_FIELD_NAME):
continue
try:
# TODO: This should support datetime and bool columns.
if promote(type_, float64, promote_option=False) != float64:
raise NotPipelineCompatible()
if isinstance(type_, Option):
@@ -307,7 +321,10 @@ def new_dataset(expr, deltas):
except TypeError:
col = NonNumpyField(name, type_)
else:
col = Column(type_.to_numpy_dtype())
col = Column(
type_.to_numpy_dtype(),
missing_values.get(name, NotSpecified),
)
columns[name] = col
@@ -473,6 +490,7 @@ def from_blaze(expr,
loader=None,
resources=None,
odo_kwargs=None,
missing_values=None,
no_deltas_rule=no_deltas_rules.warn):
"""Create a Pipeline API object from a blaze expression.
@@ -494,6 +512,9 @@ def from_blaze(expr,
scope for ``bz.compute``.
odo_kwargs : dict, optional
The keyword arguments to pass to odo when evaluating the expressions.
missing_values : dict[str -> any], optional
A dict mapping column names to missing values for those columns.
Missing values are required for integral columns.
no_deltas_rule : no_deltas_rule
What should happen if ``deltas='auto'`` but no deltas can be found.
'warn' says to raise a warning but continue.
@@ -583,7 +604,10 @@ def from_blaze(expr,
_check_resources('deltas', deltas, resources)
# Create or retrieve the Pipeline API dataset.
ds = new_dataset(dataset_expr, deltas)
if missing_values is None:
missing_values = {}
ds = new_dataset(dataset_expr, deltas, frozenset(missing_values.items()))
# Register our new dataset with the loader.
(loader if loader is not None else global_loader)[ds] = ExprData(
bind_expression_to_resources(dataset_expr, resources),
@@ -1018,7 +1042,8 @@ class BlazeLoader(dict):
column_name,
asset_idx,
sparse_deltas,
)
),
column.missing_value,
)
global_loader = BlazeLoader.global_instance()
@@ -81,12 +81,16 @@ class USEquityPricingLoader(PipelineLoader):
dates,
assets,
)
adjusted_arrays = [
AdjustedArray(raw_array, mask, col_adjustments)
for raw_array, col_adjustments in zip(raw_arrays, adjustments)
]
return dict(zip(columns, adjusted_arrays))
out = {}
for c, c_raw, c_adjs in zip(columns, raw_arrays, adjustments):
out[c] = AdjustedArray(
c_raw.astype(c.dtype),
mask,
c_adjs,
c.missing_value,
)
return out
def _shift_dates(dates, start_date, end_date, shift):
+2 -1
View File
@@ -60,7 +60,7 @@ class DataFrameLoader(PipelineLoader):
def __init__(self, column, baseline, adjustments=None):
self.column = column
self.baseline = baseline.values
self.baseline = baseline.values.astype(self.column.dtype)
self.dates = baseline.index
self.assets = baseline.columns
@@ -171,5 +171,6 @@ class DataFrameLoader(PipelineLoader):
# Mask out requested columns/rows that didnt match.
mask=(good_assets & good_dates[:, None]) & mask,
adjustments=self.format_adjustments(dates, assets),
missing_value=column.missing_value,
),
}
+2
View File
@@ -47,6 +47,7 @@ class CustomTermMixin(object):
inputs=NotSpecified,
window_length=NotSpecified,
dtype=NotSpecified,
missing_value=NotSpecified,
**kwargs):
unexpected_keys = set(kwargs) - set(cls.params)
@@ -64,6 +65,7 @@ class CustomTermMixin(object):
inputs=inputs,
window_length=window_length,
dtype=dtype,
missing_value=missing_value,
**kwargs
)
+26 -17
View File
@@ -15,7 +15,10 @@ from zipline.errors import (
WindowLengthNotSpecified,
)
from zipline.utils.memoize import lazyval
from zipline.utils.numpy_utils import bool_dtype, default_fillvalue_for_dtype
from zipline.utils.numpy_utils import (
bool_dtype,
default_missing_value_for_dtype,
)
from zipline.utils.sentinel import sentinel
@@ -32,6 +35,7 @@ class Term(with_metaclass(ABCMeta, object)):
# These are NotSpecified because a subclass is required to provide them.
dtype = NotSpecified
domain = NotSpecified
missing_value = NotSpecified
# Subclasses aren't required to provide `params`. The default behavior is
# no params.
@@ -42,6 +46,7 @@ class Term(with_metaclass(ABCMeta, object)):
def __new__(cls,
domain=domain,
dtype=dtype,
missing_value=missing_value,
# params is explicitly not allowed to be passed to an instance.
*args,
**kwargs):
@@ -55,18 +60,22 @@ class Term(with_metaclass(ABCMeta, object)):
Caching previously-constructed Terms is **sane** because terms and
their inputs are both conceptually immutable.
"""
# Class-level attributes can be used to provide defaults for Term
# subclasses.
# Subclasses can set override these class-level attributes to provide
# default values.
if domain is NotSpecified:
domain = cls.domain
if dtype is NotSpecified:
dtype = cls.dtype
if missing_value is NotSpecified:
missing_value = cls.missing_value
dtype = cls._validate_dtype(dtype)
dtype, missing_value = cls._validate_dtype(dtype, missing_value)
params = cls._pop_params(kwargs)
identity = cls.static_identity(
domain=domain,
dtype=dtype,
missing_value=missing_value,
params=params,
*args, **kwargs
)
@@ -78,6 +87,7 @@ class Term(with_metaclass(ABCMeta, object)):
super(Term, cls).__new__(cls)._init(
domain=domain,
dtype=dtype,
missing_value=missing_value,
params=params,
*args, **kwargs
)
@@ -132,9 +142,9 @@ class Term(with_metaclass(ABCMeta, object)):
return tuple(zip(cls.params, param_values))
@classmethod
def _validate_dtype(cls, passed_dtype):
def _validate_dtype(cls, passed_dtype, missing_value):
"""
Validate a `dtype` passed to Term.__new__.
Validate `dtype` passed to Term.__new__.
If passed_dtype is NotSpecified, then we try to fall back to a
class-level attribute. If a value is found at that point, we pass it
@@ -156,15 +166,17 @@ class Term(with_metaclass(ABCMeta, object)):
coercible to a numpy dtype.
"""
dtype = passed_dtype
if dtype is NotSpecified:
dtype = cls.dtype
if dtype is NotSpecified:
raise DTypeNotSpecified(termname=cls.__name__)
try:
dtype = dtype_class(dtype)
except TypeError:
raise InvalidDType(dtype=dtype, termname=cls.__name__)
return dtype
if missing_value is NotSpecified:
missing_value = default_missing_value_for_dtype(dtype)
return dtype, missing_value
def __init__(self, *args, **kwargs):
"""
@@ -183,7 +195,7 @@ class Term(with_metaclass(ABCMeta, object)):
pass
@classmethod
def static_identity(cls, domain, dtype, params):
def static_identity(cls, domain, dtype, missing_value, params):
"""
Return the identity of the Term that would be constructed from the
given arguments.
@@ -195,9 +207,9 @@ class Term(with_metaclass(ABCMeta, object)):
This is a classmethod so that it can be called from Term.__new__ to
determine whether to produce a new instance.
"""
return (cls, domain, dtype, params)
return (cls, domain, dtype, missing_value, params)
def _init(self, domain, dtype, params):
def _init(self, domain, dtype, missing_value, params):
"""
Parameters
----------
@@ -210,6 +222,7 @@ class Term(with_metaclass(ABCMeta, object)):
"""
self.domain = domain
self.dtype = dtype
self.missing_value = missing_value
for name, value in params:
if hasattr(self, name):
@@ -268,10 +281,6 @@ class Term(with_metaclass(ABCMeta, object)):
return not any(dep for dep in self.dependencies
if dep is not AssetExists())
@lazyval
def missing_value(self):
return default_fillvalue_for_dtype(self.dtype)
class AssetExists(Term):
"""
+16 -2
View File
@@ -16,7 +16,10 @@ from toolz import flip
uint8_dtype = dtype('uint8')
bool_dtype = dtype('bool')
int64_dtype = dtype('int64')
float32_dtype = dtype('float32')
float64_dtype = dtype('float64')
datetime64D_dtype = dtype('datetime64[D]')
datetime64ns_dtype = dtype('datetime64[ns]')
@@ -33,16 +36,27 @@ NaTD = NaT_for_dtype(datetime64D_dtype)
_FILLVALUE_DEFAULTS = {
bool_dtype: False,
float32_dtype: nan,
float64_dtype: nan,
datetime64ns_dtype: NaTns,
}
def default_fillvalue_for_dtype(dtype):
class NoDefaultMissingValue(Exception):
pass
def default_missing_value_for_dtype(dtype):
"""
Get the default fill value for `dtype`.
"""
return _FILLVALUE_DEFAULTS[dtype]
try:
return _FILLVALUE_DEFAULTS[dtype]
except KeyError:
raise NoDefaultMissingValue(
"No default value registered for dtype %s." % dtype
)
def repeat_first_axis(array, count):