mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-04 12:45:06 +08:00
API: DataFrame/Panel sources expect integer sids, not identifiers
This commit modifies the DataFrameSource and DataPanelSource to accept only Int64Indexes on the incoming data and moves the burden of mapping user identifiers to TradingAlgorithm.run().
This commit is contained in:
@@ -718,3 +718,25 @@ class AssetFinderTestCase(TestCase):
|
||||
|
||||
# No contracts exist after 12/14/2015, so we should get none
|
||||
self.assertIsNone(finder.lookup_future_by_expiration('AD', dt, jan_16))
|
||||
|
||||
def test_map_identifier_list_to_sids(self):
|
||||
|
||||
# Build an empty finder and some Assets
|
||||
dt = pd.Timestamp('2014-01-01', tz='UTC')
|
||||
finder = AssetFinder()
|
||||
asset1 = Equity(1, symbol="AAPL")
|
||||
asset2 = Equity(2, symbol="GOOG")
|
||||
asset200 = Future(200, symbol="CLK15")
|
||||
asset201 = Future(201, symbol="CLM15")
|
||||
|
||||
# Check for correct mapping and types
|
||||
pre_map = [asset1, asset2, asset200, asset201]
|
||||
post_map = finder.map_identifier_list_to_sids(pre_map, dt)
|
||||
self.assertListEqual([1, 2, 200, 201], post_map)
|
||||
for sid in post_map:
|
||||
self.assertIsInstance(sid, int)
|
||||
|
||||
# Change order and check mapping again
|
||||
pre_map = [asset201, asset2, asset200, asset1]
|
||||
post_map = finder.map_identifier_list_to_sids(pre_map, dt)
|
||||
self.assertListEqual([201, 2, 200, 1], post_map)
|
||||
|
||||
@@ -437,7 +437,7 @@ def handle_data(context, data):
|
||||
|
||||
_, df = factory.create_test_df_source(sim_params)
|
||||
df = df.astype(np.float64)
|
||||
source = DataFrameSource(df, sids=[0])
|
||||
source = DataFrameSource(df)
|
||||
|
||||
test_algo = TradingAlgorithm(
|
||||
script=algo_text,
|
||||
|
||||
+11
-5
@@ -27,6 +27,7 @@ from zipline.sources import (DataFrameSource,
|
||||
RandomWalkSource)
|
||||
from zipline.utils import tradingcalendar as calendar_nyse
|
||||
from zipline.finance.trading import with_environment
|
||||
from zipline.assets import AssetFinder
|
||||
|
||||
|
||||
class TestDataFrameSource(TestCase):
|
||||
@@ -43,7 +44,7 @@ class TestDataFrameSource(TestCase):
|
||||
|
||||
def test_df_sid_filtering(self):
|
||||
_, df = factory.create_test_df_source()
|
||||
source = DataFrameSource(df, sids=[0])
|
||||
source = DataFrameSource(df)
|
||||
assert 1 not in [event.sid for event in source], \
|
||||
"DataFrameSource should only stream selected sid 0, not sid 1."
|
||||
|
||||
@@ -63,8 +64,8 @@ class TestDataFrameSource(TestCase):
|
||||
self.assertTrue(isinstance(event['volume'], int))
|
||||
self.assertTrue(isinstance(event['arbitrary'], float))
|
||||
|
||||
@with_environment()
|
||||
def test_yahoo_bars_to_panel_source(self, env=None):
|
||||
def test_yahoo_bars_to_panel_source(self):
|
||||
finder = AssetFinder()
|
||||
stocks = ['AAPL', 'GE']
|
||||
start = pd.datetime(1993, 1, 1, 0, 0, 0, 0, pytz.utc)
|
||||
end = pd.datetime(2002, 1, 1, 0, 0, 0, 0, pytz.utc)
|
||||
@@ -75,10 +76,15 @@ class TestDataFrameSource(TestCase):
|
||||
|
||||
check_fields = ['sid', 'open', 'high', 'low', 'close',
|
||||
'volume', 'price']
|
||||
source = DataPanelSource(data)
|
||||
|
||||
copy_panel = data.copy()
|
||||
copy_panel.items = finder.map_identifier_list_to_sids(
|
||||
data.items, data.major_axis[0]
|
||||
)
|
||||
source = DataPanelSource(copy_panel)
|
||||
sids = [
|
||||
asset.sid for asset in
|
||||
[env.asset_finder.lookup_symbol(symbol, as_of_date=end)
|
||||
[finder.lookup_symbol(symbol, as_of_date=end)
|
||||
for symbol in stocks]
|
||||
]
|
||||
stocks_iter = cycle(sids)
|
||||
|
||||
@@ -98,7 +98,7 @@ class TestTALIB(TestCase):
|
||||
zipline_transforms = [ta.MA(timeperiod=10),
|
||||
ta.MA(timeperiod=25)]
|
||||
talib_fn = talib.abstract.MA
|
||||
algo = TALIBAlgorithm(talib=zipline_transforms)
|
||||
algo = TALIBAlgorithm(talib=zipline_transforms, identifiers=[0])
|
||||
algo.run(self.source)
|
||||
# Test if computed values match those computed by pandas rolling mean.
|
||||
sid = 0
|
||||
|
||||
Reference in New Issue
Block a user