Converts data frame source over to using DataSource base object.

This commit is contained in:
Eddie Hebert
2012-11-19 17:30:16 -05:00
parent f85ad50e60
commit 5a95122f21
2 changed files with 94 additions and 31 deletions
+30 -31
View File
@@ -16,20 +16,14 @@
"""
Tools to generate data sources.
"""
from copy import copy
from itertools import ifilter
import pandas as pd
from zipline.gens.utils import hash_args
from zipline.protocol import DATASOURCE_TYPE
from zipline.utils import ndict
from zipline.sources.test_source import SpecificEquityTrades
from zipline.sources.data_source import DataSource
class DataFrameSource(SpecificEquityTrades):
class DataFrameSource(DataSource):
"""
Yields all events in event_list that match the given sid_filter.
If no event_list is specified, generates an internal stream of events
@@ -37,7 +31,6 @@ class DataFrameSource(SpecificEquityTrades):
Configuration options:
count : integer representing number of trades
sids : list of values representing simulated internal sids
start : start date
delta : timedelta between internal events
@@ -49,36 +42,42 @@ class DataFrameSource(SpecificEquityTrades):
self.data = data
# Unpack config dictionary with default values.
self.count = kwargs.get('count', len(data))
self.sids = kwargs.get('sids', data.columns)
self.start = kwargs.get('start', data.index[0])
self.end = kwargs.get('end', data.index[-1])
self.delta = kwargs.get('delta', data.index[1] - data.index[0])
# Hash_value for downstream sorting.
self.arg_string = hash_args(data, **kwargs)
self.generator = self.create_fresh_generator()
self._raw_data = None
def create_fresh_generator(self):
def _generator(df=self.data):
for dt, series in df.iterrows():
if (dt < self.start) or (dt > self.end):
continue
event = {
'dt': dt,
'source_id': self.get_hash(),
'type': DATASOURCE_TYPE.TRADE
}
@property
def mapping(self):
return {
'dt': (lambda x: x, 'dt'),
'sid': (lambda x: x, 'sid'),
'price': (float, 'price'),
'volume': (int, 'volume'),
}
for sid, price in series.iterkv():
event = copy(event)
event['sid'] = sid
event['price'] = price
event['volume'] = 1000
@property
def instance_hash(self):
return self.arg_string
yield ndict(event)
def raw_data_gen(self):
for dt, series in self.data.iterrows():
for sid, price in series.iterkv():
if sid in self.sids:
event = {
'dt': dt,
'sid': sid,
'price': price,
'volume': 1000,
}
yield event
# Return the filtered event stream.
drop_sids = lambda x: x.sid in self.sids
return ifilter(drop_sids, _generator())
@property
def raw_data(self):
if not self._raw_data:
self._raw_data = self.raw_data_gen()
return self._raw_data
+64
View File
@@ -0,0 +1,64 @@
from abc import (
ABCMeta,
abstractproperty
)
from zipline.protocol import DATASOURCE_TYPE
from zipline.utils.protocol_utils import ndict
class DataSource(object):
__metaclass__ = ABCMeta
@property
def event_type(self):
return DATASOURCE_TYPE.TRADE
@property
def mapping(self):
"""
Mappings of the form:
target_key: (mapping_function, source_key)
"""
return {}
@abstractproperty
def raw_data(self):
"""
An iterator that yields the raw datasource,
in chronological order of data, one event at a time.
"""
NotImplemented
@abstractproperty
def instance_hash(self):
"""
A hash that represents the unique args to the source.
"""
pass
def get_hash(self):
return self.__class__.__name__ + "-" + self.instance_hash
def apply_mapping(self, raw_row):
"""
Override this to hand craft conversion of row.
"""
row = {target: mapping_func(raw_row[source_key])
for target, (mapping_func, source_key)
in self.mapping.items()}
row.update({'source_id': self.get_hash()})
row.update({'type': self.event_type})
return row
@property
def mapped_data(self):
for row in self.raw_data:
yield ndict(self.apply_mapping(row))
def __iter__(self):
return self
def next(self):
return self.mapped_data.next()