mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-10 11:50:32 +08:00
add groupby to rank, top, and bottom
This commit is contained in:
committed by
Scott Sanderson
parent
161922897e
commit
9e3404646e
@@ -6,6 +6,7 @@ from operator import attrgetter
|
||||
from numbers import Number
|
||||
|
||||
from numpy import inf, where
|
||||
from scipy.stats import rankdata
|
||||
|
||||
from zipline.errors import UnknownRankMethod
|
||||
from zipline.lib.normalize import naive_grouped_rowwise_apply
|
||||
@@ -581,7 +582,11 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
window_safe=True,
|
||||
)
|
||||
|
||||
def rank(self, method='ordinal', ascending=True, mask=NotSpecified):
|
||||
def rank(self,
|
||||
method='ordinal',
|
||||
ascending=True,
|
||||
mask=NotSpecified,
|
||||
groupby=NotSpecified):
|
||||
"""
|
||||
Construct a new Factor representing the sorted rank of each column
|
||||
within each row.
|
||||
@@ -599,6 +604,8 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
A Filter representing assets to consider when computing ranks.
|
||||
If mask is supplied, ranks are computed ignoring any asset/date
|
||||
pairs for which `mask` produces a value of False.
|
||||
groupby : zipline.pipeline.Classifier, optional
|
||||
A classifier defining partitions over which to perform ranking.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -620,7 +627,21 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
:func:`scipy.stats.rankdata`
|
||||
:class:`zipline.pipeline.factors.factor.Rank`
|
||||
"""
|
||||
return Rank(self, method=method, ascending=ascending, mask=mask)
|
||||
|
||||
if groupby is NotSpecified:
|
||||
return Rank(self, method=method, ascending=ascending, mask=mask)
|
||||
|
||||
else:
|
||||
def rank(row):
|
||||
return rankdata(row if ascending else -row, method=method)
|
||||
|
||||
return GroupedRowTransform(
|
||||
transform=rank,
|
||||
factor=self,
|
||||
mask=mask,
|
||||
groupby=groupby,
|
||||
window_safe=True,
|
||||
)
|
||||
|
||||
@expect_types(
|
||||
target=Term, correlation_length=int, mask=(Filter, NotSpecifiedType),
|
||||
@@ -913,7 +934,7 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
"""
|
||||
return self.quantiles(bins=10, mask=mask)
|
||||
|
||||
def top(self, N, mask=NotSpecified):
|
||||
def top(self, N, mask=NotSpecified, groupby=NotSpecified):
|
||||
"""
|
||||
Construct a Filter matching the top N asset values of self each day.
|
||||
|
||||
@@ -925,14 +946,16 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
A Filter representing assets to consider when computing ranks.
|
||||
If mask is supplied, top values are computed ignoring any
|
||||
asset/date pairs for which `mask` produces a value of False.
|
||||
groupby : zipline.pipeline.Classifier, optional
|
||||
A classifier defining partitions over which to perform ranking.
|
||||
|
||||
Returns
|
||||
-------
|
||||
filter : zipline.pipeline.filters.Filter
|
||||
"""
|
||||
return self.rank(ascending=False, mask=mask) <= N
|
||||
return self.rank(ascending=False, mask=mask, groupby=groupby) <= N
|
||||
|
||||
def bottom(self, N, mask=NotSpecified):
|
||||
def bottom(self, N, mask=NotSpecified, groupby=NotSpecified):
|
||||
"""
|
||||
Construct a Filter matching the bottom N asset values of self each day.
|
||||
|
||||
@@ -944,12 +967,14 @@ class Factor(RestrictedDTypeMixin, ComputableTerm):
|
||||
A Filter representing assets to consider when computing ranks.
|
||||
If mask is supplied, bottom values are computed ignoring any
|
||||
asset/date pairs for which `mask` produces a value of False.
|
||||
groupby : zipline.pipeline.Classifier, optional
|
||||
A classifier defining partitions over which to perform ranking.
|
||||
|
||||
Returns
|
||||
-------
|
||||
filter : zipline.pipeline.Filter
|
||||
"""
|
||||
return self.rank(ascending=True, mask=mask) <= N
|
||||
return self.rank(ascending=True, mask=mask, groupby=groupby) <= N
|
||||
|
||||
def percentile_between(self,
|
||||
min_percentile,
|
||||
@@ -1075,7 +1100,7 @@ class GroupedRowTransform(Factor):
|
||||
Factor.
|
||||
|
||||
This is most often useful for normalization operators like ``zscore`` or
|
||||
``demean``.
|
||||
``demean`` or for performing ranking using ``rank``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
@@ -1093,12 +1118,13 @@ class GroupedRowTransform(Factor):
|
||||
-----
|
||||
Users should rarely construct instances of this factor directly. Instead,
|
||||
they should construct instances via factor normalization methods like
|
||||
``zscore`` and ``demean``.
|
||||
``zscore`` and ``demean`` or using ``rank`` with ``groupby``.
|
||||
|
||||
See Also
|
||||
--------
|
||||
zipline.pipeline.factors.Factor.zscore
|
||||
zipline.pipeline.factors.Factor.demean
|
||||
zipline.pipeline.factors.Factor.rank
|
||||
"""
|
||||
window_length = 0
|
||||
|
||||
|
||||
Reference in New Issue
Block a user