mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-09 11:19:23 +08:00
Merge pull request #1195 from quantopian/test-string-groupby
String Classifier Cleanup
This commit is contained in:
@@ -8,7 +8,6 @@ from zipline.pipeline import Classifier
|
||||
from zipline.testing import parameter_space
|
||||
from zipline.utils.numpy_utils import (
|
||||
categorical_dtype,
|
||||
coerce_to_dtype,
|
||||
int64_dtype,
|
||||
)
|
||||
|
||||
@@ -162,7 +161,8 @@ class ClassifierTestCase(BasePipelineTestCase):
|
||||
dtype_=[int64_dtype, categorical_dtype],
|
||||
)
|
||||
def test_disallow_comparison_to_missing_value(self, missing, dtype_):
|
||||
missing = coerce_to_dtype(dtype_, missing)
|
||||
if dtype_ == categorical_dtype:
|
||||
missing = str(missing)
|
||||
|
||||
class C(Classifier):
|
||||
dtype = dtype_
|
||||
@@ -434,7 +434,7 @@ class ClassifierTestCase(BasePipelineTestCase):
|
||||
errmsg = str(e.exception)
|
||||
expected = (
|
||||
"Found self.missing_value ('not in the array') in choices"
|
||||
" supplied to C.is_element().\n"
|
||||
" supplied to C.element_of().\n"
|
||||
"Missing values have NaN semantics, so the requested"
|
||||
" comparison would always produce False.\n"
|
||||
"Use the isnull() method to check for missing values.\n"
|
||||
@@ -447,7 +447,7 @@ class ClassifierTestCase(BasePipelineTestCase):
|
||||
|
||||
class C(Classifier):
|
||||
dtype = dtype_
|
||||
missing_value = ''
|
||||
missing_value = dtype.type('1')
|
||||
inputs = ()
|
||||
window_length = 0
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from numpy import (
|
||||
from numpy.random import randn, seed
|
||||
|
||||
from zipline.errors import UnknownRankMethod
|
||||
from zipline.lib.labelarray import LabelArray
|
||||
from zipline.lib.rank import masked_rankdata_2d
|
||||
from zipline.lib.normalize import naive_grouped_rowwise_apply as grouped_apply
|
||||
from zipline.pipeline import Classifier, Factor, Filter, TermGraph
|
||||
@@ -38,6 +39,7 @@ from zipline.testing import (
|
||||
)
|
||||
from zipline.utils.functional import dzip_exact
|
||||
from zipline.utils.numpy_utils import (
|
||||
categorical_dtype,
|
||||
datetime64ns_dtype,
|
||||
float64_dtype,
|
||||
int64_dtype,
|
||||
@@ -442,6 +444,7 @@ class FactorTestCase(BasePipelineTestCase):
|
||||
f = self.f
|
||||
m = Mask()
|
||||
c = C()
|
||||
str_c = C(dtype=categorical_dtype, missing_value=None)
|
||||
|
||||
factor_data = array(
|
||||
[[1.0, 2.0, 3.0, 4.0],
|
||||
@@ -463,12 +466,18 @@ class FactorTestCase(BasePipelineTestCase):
|
||||
[1, 1, 2, 2]],
|
||||
dtype=int64_dtype,
|
||||
)
|
||||
string_classifier_data = LabelArray(
|
||||
classifier_data.astype(str).astype(object),
|
||||
missing_value=None,
|
||||
)
|
||||
|
||||
terms = {
|
||||
'vanilla': f.demean(),
|
||||
'masked': f.demean(mask=m),
|
||||
'grouped': f.demean(groupby=c),
|
||||
'grouped_str': f.demean(groupby=str_c),
|
||||
'grouped_masked': f.demean(mask=m, groupby=c),
|
||||
'grouped_masked_str': f.demean(mask=m, groupby=str_c),
|
||||
}
|
||||
expected = {
|
||||
'vanilla': array(
|
||||
@@ -496,6 +505,9 @@ class FactorTestCase(BasePipelineTestCase):
|
||||
[-0.500, 0.500, 0.000, nan]]
|
||||
)
|
||||
}
|
||||
# Changing the classifier dtype shouldn't affect anything.
|
||||
expected['grouped_str'] = expected['grouped']
|
||||
expected['grouped_masked_str'] = expected['grouped_masked']
|
||||
|
||||
graph = TermGraph(terms)
|
||||
results = self.run_graph(
|
||||
@@ -503,6 +515,7 @@ class FactorTestCase(BasePipelineTestCase):
|
||||
initial_workspace={
|
||||
f: factor_data,
|
||||
c: classifier_data,
|
||||
str_c: string_classifier_data,
|
||||
m: filter_data,
|
||||
},
|
||||
mask=self.build_mask(self.ones_mask(shape=factor_data.shape)),
|
||||
|
||||
@@ -602,3 +602,26 @@ class SubDataSetTestCase(TestCase):
|
||||
with self.assertRaises(ValueError) as e:
|
||||
SomeClassifier()
|
||||
self.assertEqual(str(e.exception), expected_error)
|
||||
|
||||
def test_unreasonable_missing_values(self):
|
||||
|
||||
for base_type, dtype_, bad_mv in ((Factor, float64_dtype, 'ayy'),
|
||||
(Filter, bool_dtype, 'lmao'),
|
||||
(Classifier, int64_dtype, 'lolwut'),
|
||||
(Classifier, categorical_dtype, 7)):
|
||||
class SomeTerm(base_type):
|
||||
inputs = ()
|
||||
window_length = 0
|
||||
missing_value = bad_mv
|
||||
dtype = dtype_
|
||||
|
||||
with self.assertRaises(TypeError) as e:
|
||||
SomeTerm()
|
||||
|
||||
prefix = (
|
||||
"^Missing value {mv!r} is not a valid choice "
|
||||
"for term SomeTerm with dtype {dtype}.\n\n"
|
||||
"Coercion attempt failed with:"
|
||||
).format(mv=bad_mv, dtype=dtype_)
|
||||
|
||||
self.assertRegexpMatches(str(e.exception), prefix)
|
||||
|
||||
Reference in New Issue
Block a user