Merge pull request #1665 from quantopian/determine_overwrite_type_dynamically

Determine overwrite type dynamically
This commit is contained in:
Maya Tydykov
2017-01-31 10:06:21 -05:00
committed by GitHub
5 changed files with 100 additions and 6 deletions
+3
View File
@@ -25,6 +25,7 @@ from zipline.lib.adjustment import (
Float64Multiply,
Float64Overwrite,
Float641DArrayOverwrite,
Int64Overwrite,
ObjectOverwrite,
)
from zipline.lib.adjusted_array import AdjustedArray, NOMASK
@@ -235,6 +236,7 @@ def _gen_overwrite_adjustment_cases(dtype):
adjustment_type = {
float64_dtype: Float64Overwrite,
datetime64ns_dtype: Datetime64Overwrite,
int64_dtype: Int64Overwrite,
bytes_dtype: ObjectOverwrite,
unicode_dtype: ObjectOverwrite,
object_dtype: ObjectOverwrite,
@@ -585,6 +587,7 @@ class AdjustedArrayTestCase(TestCase):
@parameterized.expand(
chain(
_gen_overwrite_adjustment_cases(int64_dtype),
_gen_overwrite_adjustment_cases(float64_dtype),
_gen_overwrite_adjustment_cases(datetime64ns_dtype),
_gen_overwrite_1d_array_adjustment_case(float64_dtype),
+32
View File
@@ -35,6 +35,21 @@ class AdjustmentTestCase(TestCase):
)
self.assertEqual(result, expected)
def test_make_int_adjustment(self):
result = adj.make_adjustment_from_indices(
1, 2, 3, 4,
adjustment_kind=adj.OVERWRITE,
value=1,
)
expected = adj.Int64Overwrite(
first_row=1,
last_row=2,
first_col=3,
last_col=4,
value=1,
)
self.assertEqual(result, expected)
def test_make_datetime_adjustment(self):
overwrite_dt = make_datetime64ns(0)
result = adj.make_adjustment_from_indices(
@@ -51,6 +66,23 @@ class AdjustmentTestCase(TestCase):
)
self.assertEqual(result, expected)
@parameterized.expand([("some text",), ("some text".encode(),), (None,)])
def test_make_object_adjustment(self, value):
result = adj.make_adjustment_from_indices(
1, 2, 3, 4,
adjustment_kind=adj.OVERWRITE,
value=value,
)
expected = adj.ObjectOverwrite(
first_row=1,
last_row=2,
first_col=3,
last_col=4,
value=value,
)
self.assertEqual(result, expected)
def test_unsupported_type(self):
class SomeClass(object):
pass
+1 -2
View File
@@ -1534,7 +1534,7 @@ class BlazeToPipelineTestCase(WithAssetFinder, ZiplineTestCase):
pd.Timestamp('2014-01-04')
])
baseline = pd.DataFrame({
'value': (0, 1),
'value': (0., 1.),
'asof_date': base_dates,
'timestamp': base_dates,
})
@@ -1545,7 +1545,6 @@ class BlazeToPipelineTestCase(WithAssetFinder, ZiplineTestCase):
value=deltas.value + 10,
timestamp=deltas.timestamp + timedelta(days=1),
)
nassets = len(simple_asset_info)
expected_views = keymap(pd.Timestamp, {
'2014-01-03': np.array([[10.0],