mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-11 11:16:15 +08:00
- Adds a new class, ``LabelArray``, which is a subclass of np.ndarray. LabelArray is conceptually similar to pandas.Categorical, in that it stores data with many duplicate values as indices into an array of unique values. For string data with many duplicates (e.g. time-series of tickers or or industry classifications), this provides multiple orders of magnitude of improvement when doing string operations, especially string comparison/matching operations. - Adds a new generic object "specialization" for `AdjustedArrayWindow`, and a corresponding ObjectOverwrite adjustment. - Adds a new ``postprocess`` method to ``zipline.pipeline.term.Term``. This method is called on the final result of any pipeline expression after screen filtering has occurred. The default implementation of ``postprocess`` is identity, but Classifier overrides it to coerce string columns into pandas.Categoricals before presenting them to the user.
233 lines
6.9 KiB
Python
233 lines
6.9 KiB
Python
"""
|
|
dataset.py
|
|
"""
|
|
from functools import total_ordering
|
|
from six import (
|
|
iteritems,
|
|
with_metaclass,
|
|
)
|
|
|
|
from zipline.pipeline.term import (
|
|
AssetExists,
|
|
LoadableTerm,
|
|
NotSpecified,
|
|
Term,
|
|
)
|
|
from zipline.utils.input_validation import ensure_dtype
|
|
from zipline.utils.numpy_utils import NoDefaultMissingValue
|
|
from zipline.utils.preprocess import preprocess
|
|
|
|
|
|
class Column(object):
|
|
"""
|
|
An abstract column of data, not yet associated with a dataset.
|
|
"""
|
|
@preprocess(dtype=ensure_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,
|
|
missing_value=self.missing_value,
|
|
name=name,
|
|
)
|
|
|
|
|
|
class _BoundColumnDescr(object):
|
|
"""
|
|
Intermediate class that sits on `DataSet` objects and returns memoized
|
|
`BoundColumn` objects when requested.
|
|
|
|
This exists so that subclasses of DataSets don't share columns with their
|
|
parent classes.
|
|
"""
|
|
def __init__(self, dtype, missing_value, name):
|
|
# Validating and calculating default missing values here guarantees
|
|
# that we fail quickly if the user passes an unsupporte dtype or 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.
|
|
try:
|
|
self.dtype, self.missing_value = Term.validate_dtype(
|
|
termname="Column(name={name!r})".format(name=name),
|
|
dtype=dtype,
|
|
missing_value=missing_value,
|
|
)
|
|
except NoDefaultMissingValue:
|
|
# Re-raise with a more specific 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.name = name
|
|
|
|
def __get__(self, instance, owner):
|
|
"""
|
|
Produce a concrete BoundColumn object when accessed.
|
|
|
|
We don't bind to datasets at class creation time so that subclasses of
|
|
DataSets produce different BoundColumns.
|
|
"""
|
|
return BoundColumn(
|
|
dtype=self.dtype,
|
|
missing_value=self.missing_value,
|
|
dataset=owner,
|
|
name=self.name,
|
|
)
|
|
|
|
|
|
class BoundColumn(LoadableTerm):
|
|
"""
|
|
A column of data that's been concretely bound to a particular dataset.
|
|
|
|
Instances of this class are dynamically created upon access to attributes
|
|
of DataSets (for example, USEquityPricing.close is an instance of this
|
|
class).
|
|
|
|
Attributes
|
|
----------
|
|
dtype : numpy.dtype
|
|
The dtype of data produced when this column is loaded.
|
|
latest : zipline.pipeline.data.Factor or zipline.pipeline.data.Filter
|
|
A Filter, Factor, or Classifier computing the most recently known value
|
|
of this column on each date.
|
|
|
|
Produces a Filter if self.dtype == ``np.bool_``.
|
|
Produces a Classifier if self.dtype == ``np.int64``
|
|
Otherwise produces a Factor.
|
|
dataset : zipline.pipeline.data.DataSet
|
|
The dataset to which this column is bound.
|
|
name : str
|
|
The name of this column.
|
|
"""
|
|
mask = AssetExists()
|
|
inputs = ()
|
|
|
|
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,
|
|
)
|
|
|
|
def _init(self, dataset, name, *args, **kwargs):
|
|
self._dataset = dataset
|
|
self._name = name
|
|
return super(BoundColumn, self)._init(*args, **kwargs)
|
|
|
|
@classmethod
|
|
def static_identity(cls, dataset, name, *args, **kwargs):
|
|
return (
|
|
super(BoundColumn, cls).static_identity(*args, **kwargs),
|
|
dataset,
|
|
name,
|
|
)
|
|
|
|
@property
|
|
def dataset(self):
|
|
"""
|
|
The dataset to which this column is bound.
|
|
"""
|
|
return self._dataset
|
|
|
|
@property
|
|
def name(self):
|
|
"""
|
|
The name of this column.
|
|
"""
|
|
return self._name
|
|
|
|
@property
|
|
def qualname(self):
|
|
"""
|
|
The fully-qualified name of this column.
|
|
|
|
Generated by doing '.'.join([self.dataset.__name__, self.name]).
|
|
"""
|
|
return '.'.join([self.dataset.__name__, self.name])
|
|
|
|
@property
|
|
def latest(self):
|
|
from zipline.pipeline.factors import Factor, Latest as LatestFactor
|
|
from zipline.pipeline.filters import Filter, Latest as LatestFilter
|
|
from zipline.pipeline.classifiers import (
|
|
Classifier,
|
|
Latest as LatestClassifier,
|
|
)
|
|
|
|
dtype = self.dtype
|
|
if dtype in Filter.ALLOWED_DTYPES:
|
|
Latest = LatestFilter
|
|
elif dtype in Classifier.ALLOWED_DTYPES:
|
|
Latest = LatestClassifier
|
|
else:
|
|
assert dtype in Factor.ALLOWED_DTYPES, "Unknown dtype %s." % dtype
|
|
Latest = LatestFactor
|
|
|
|
return Latest(
|
|
inputs=(self,),
|
|
dtype=dtype,
|
|
missing_value=self.missing_value,
|
|
)
|
|
|
|
def __repr__(self):
|
|
return "{qualname}::{dtype}".format(
|
|
qualname=self.qualname,
|
|
dtype=self.dtype.name,
|
|
)
|
|
|
|
def short_repr(self):
|
|
return self.qualname
|
|
|
|
|
|
@total_ordering
|
|
class DataSetMeta(type):
|
|
"""
|
|
Metaclass for DataSets
|
|
|
|
Supplies name and dataset information to Column attributes.
|
|
"""
|
|
|
|
def __new__(mcls, name, bases, dict_):
|
|
newtype = super(DataSetMeta, mcls).__new__(mcls, name, bases, dict_)
|
|
# collect all of the column names that we inherit from our parents
|
|
column_names = set().union(
|
|
*(getattr(base, '_column_names', ()) for base in bases)
|
|
)
|
|
for maybe_colname, maybe_column in iteritems(dict_):
|
|
if isinstance(maybe_column, Column):
|
|
# add column names defined on our class
|
|
bound_column_descr = maybe_column.bind(maybe_colname)
|
|
setattr(newtype, maybe_colname, bound_column_descr)
|
|
column_names.add(maybe_colname)
|
|
|
|
newtype._column_names = frozenset(column_names)
|
|
return newtype
|
|
|
|
@property
|
|
def columns(self):
|
|
return frozenset(
|
|
getattr(self, colname) for colname in self._column_names
|
|
)
|
|
|
|
def __lt__(self, other):
|
|
return id(self) < id(other)
|
|
|
|
def __repr__(self):
|
|
return '<DataSet: %r>' % self.__name__
|
|
|
|
|
|
class DataSet(with_metaclass(DataSetMeta, object)):
|
|
domain = None
|