Merge branch 'byebye_threadsim' of github.com:quantopian/zipline into byebye_threadsim

This commit is contained in:
fawce
2012-07-27 17:15:05 -04:00
7 changed files with 265 additions and 36 deletions
+3 -6
View File
@@ -12,7 +12,7 @@ from datetime import datetime, timedelta
from collections import deque, defaultdict
from zipline import ndict
from zipline.gens.utils import stringify_args, assert_datasource_protocol, \
from zipline.gens.utils import hash_args, assert_datasource_protocol, \
assert_trade_protocol, assert_datasource_unframe_protocol
import zipline.protocol as zp
@@ -25,7 +25,6 @@ def FeedGen(stream_in, source_ids):
message and yield it.
"""
assert isinstance(stream_in, types.GeneratorType)
assert isinstance(source_ids, list)
# Set up an internal queue for each expected source.
@@ -34,10 +33,7 @@ def FeedGen(stream_in, source_ids):
assert isinstance(id, basestring), "Bad source_id %s" % source_id
sources[id] = deque()
namestring = "FeedGen" + stringify_args(source_ids)
# Process incoming streams.
for message in stream_in:
# Incoming messages should be the output of DATASOURCE_UNFRAME.
assert_datasource_unframe_protocol(message), \
@@ -53,6 +49,7 @@ def FeedGen(stream_in, source_ids):
while full(sources) and not done(sources):
message = pop_oldest(sources)
assert feed_protocol(message)
yield message
# We should have only a done message left in each queue.
@@ -109,7 +106,7 @@ def pop_oldest(sources):
oldest_event = older(oldest_event, current_event)
# Pop the oldest event we found from its queue and return it.
return sources[oldest_event.source_id].pop()
return sources[oldest_event.source_id].popleft()
# Return the event with the older timestamp. Break ties by source_id.
def older(oldest, current):
+2 -2
View File
@@ -10,7 +10,7 @@ from pymongo import ASCENDING
from datetime import datetime, timedelta
from zipline import ndict
from zipline.gens.utils import stringify_args, assert_datasource_protocol, \
from zipline.gens.utils import hash_args, assert_datasource_protocol, \
assert_trade_protocol
import zipline.protocol as zp
@@ -33,7 +33,7 @@ def MongoTradeHistoryGen(collection, filter, start_date, end_date):
# Create unique identifier string that can be used to break
# sorting ties deterministically
argstring = stringify_args(collection, filter, start_date, end_date)
argstring = hash_args(collection, filter, start_date, end_date)
source_id = "MongoTradeHistoryGen" + argstring
# All datasources
+20 -16
View File
@@ -6,7 +6,7 @@ import pytz
from unittest2 import TestCase
from pymongo import Connection, ASCENDING
from itertools import izip, izip_longest, permutations, cycle
from itertools import izip, izip_longest, permutations, cycle, chain
from datetime import datetime, timedelta
from collections import deque
@@ -14,7 +14,7 @@ from zipline import ndict
from zipline.gens.feed import FeedGen, full, done, queue_is_full,queue_is_done,\
pop_oldest
from zipline.gens.utils import stringify_args, assert_datasource_protocol,\
assert_trade_protocol, date_gen
assert_trade_protocol, date_gen, alternate
import zipline.protocol as zp
@@ -96,7 +96,8 @@ class FeedGenTestCase(TestCase):
Assert that FeedGen's output agrees with expected.
"""
feed_gen = FeedGen(events, source_ids)
assert list(feed_gen) == expected
l = list(feed_gen)
assert l == expected
def test_single_source(self):
@@ -119,38 +120,41 @@ class FeedGenTestCase(TestCase):
self.run_FeedGen(event_gen, expected, source_ids)
def test_multi_source_interleaved(self):
def test_multi_source(self):
source_ids = ['a', 'b']
type = zp.DATASOURCE_TYPE.TRADE
# Set up source 'a'. Outputs 3 events with 2 minute deltas.
# Set up source 'a'. Outputs 20 events with 2 minute deltas.
delta_a = timedelta(minutes = 2)
dates_a = list(date_gen(delta = delta_a, n = 3))
dates_a = list(date_gen(delta = delta_a, n = 20))
dates_a.append("DONE")
events_a_args = zip(cycle(['a']), iter(dates_a), cycle([type]))
events_a = [mock_data_unframe(*args) for args in events_a_args]
event_gen_a = (e for e in events_a)
# Set up source 'b'. Outputs 4 events with 1 minute deltas.
# Set up source 'b'. Outputs 10 events with 1 minute deltas.
delta_b = timedelta(minutes = 1)
dates_b = list(date_gen(delta = delta_b, n = 4))
dates_b = list(date_gen(delta = delta_b, n = 10))
dates_b.append("DONE")
events_b_args = zip(cycle(['b']), iter(dates_b), cycle([type]))
events_b = [mock_data_unframe(*args) for args in events_b_args]
event_gen_b = (e for e in events_b)
# The expected output is all non-DONE events in both a and b,
# sorted first by dt and then by source_id.
non_dones = events_a[:-1] + events_b[:-1]
expected = sorted(non_dones, compare_by_dt_source_id)
import nose.tools; nose.tools.set_trace()
self.run_FeedGen(event_gen, expected, source_ids)
# Alternating between a and b.
interleaved = alternate(iter(events_a), iter(events_b))
self.run_FeedGen(interleaved, expected, source_ids)
# All of a, then all of b.
sequential = chain(iter(events_a), iter(events_b))
self.run_FeedGen(sequential, expected, source_ids)
# def test_FeedGen_consistency(self):
# source_ids = ['a', 'b']
-9
View File
@@ -15,15 +15,6 @@ from zipline.gens.utils import stringify_args, assert_datasource_protocol,\
import zipline.protocol as zp
def mock_raw_event(sid, dt):
event = {
'sid' : sid,
'dt' : dt,
'price' : 1.0,
'volume' : 1
}
return event
mongo_conn_args = {
'mongodb_host': 'localhost',
'mongodb_port': 27017,
+211
View File
@@ -0,0 +1,211 @@
"""
Generator versions of transforms.
"""
import random
import pytz
import logbook
import pymongo
import types
from pymongo import ASCENDING
from datetime import datetime, timedelta
from collections import deque, defaultdict
from numbers import Number
from itertools import izip
from zipline import ndict
from zipline.gens.utils import hash_args, date_gen
from zipline.gens.utils import assert_feed_unframe_protocol, assert_transform_protocol
import zipline.protocol as zp
def PassthroughTransformGen(stream_in):
"""Trivial transform for event forwarding."""
# hash_args with no arguments is the same as:
# hasher = hashlib.md5()
# hasher.update(":");
# hashlib.md5.digest().
namestring = "Passthrough" + hash_args()
for message in stream_in:
assert_feed_unframe_protocol(message)
out_value = message
assert_transform_protocol(out_value)
yield (namestring, out_value)
def FunctionalTransformGen(stream_in, fun, *args, **kwargs):
"""
Generic transform generator that takes each message from an in-stream
and yields the output of a function on that message. Not sure how
useful this will be in reality, but good for testing.
"""
# TODO: Distinguish between functions and classes in hash_args.
namestring = fun.__name__ + hash_args(*args, **kwargs)
for message in stream_in:
assert_feed_unframe_protocol(message)
out_value = fun(message, *args, **kwargs)
assert_transform_protocol(out_value)
yield(namestring, out_value)
def StatefulTransformGen(stream_in, tnfm_class, *args, **kwargs):
"""
Generic transform generator that takes each message from an in-stream
and feeds it to a state class. For each call to update, the state
class must produce a message to be fed downstream.
"""
# Create an instance of our transform class.
state = tnfm_class(*args, **kwargs)
# Generate the string associated with this generator's output.
namestring = tnfm_class.__name__ + hash_args(*args, **kwargs)
for message in stream_in:
assert_feed_unframe_protocol(message)
out_value = state.update(message)
assert_transform_protocol(out_value)
yield (namestring, out_value)
def MovingAverageTransformGen(stream_in, days, fields):
"""
Generator that uses the MovingAverage state class to calculate
a moving average for all stocks over a specified number of days.
"""
return StatefulTransformGen(stream_in, MovingAverage, timedelta(days=days), fields)
class MovingAverage(object):
"""
Class that maintains a dictionary from sids to EventWindows
Upon receipt of each message we update the
corresponding window and return the calculated average.
"""
def __init__(self, delta, fields):
self.delta = delta
self.fields = fields
# No way to pass arguments to the defaultdict factory, so we
# need to define a method to generate the correct EventWindows.
self.sid_windows = defaultdict(self.create_window)
def create_window(self):
"""Factory method for self.sid_windows."""
return EventWindow(self.delta, self.fields)
def update(self, event):
"""
Update the event window for this event's sid. Return an ndict from
tracked fields to averages.
"""
assert isinstance(event, ndict),"Bad event in MovingAverage: %s" % event
assert event.has_key('sid'), "No sid in MovingAverage: %s" % event
# This will create a new EventWindow if this is the first
# message for this sid.
window = self.sid_windows[event.sid]
window.update(event)
return window.get_averages()
class EventWindow(object):
"""
Maintains a list of events that are within a certain timedelta
of the most recent tick. The expected use of this class is to
track events associated with a single sid. We provide simple
functionality for averages, but anything more complicated
should be handled by a containing class.
"""
def __init__(self, delta, fields):
self.ticks = deque()
self.delta = delta
self.fields = fields
self.totals = defaultdict(float)
def __len__(self):
return len(self.ticks)
def update(self, event):
self.assert_well_formed(event)
# Add new event and increment totals.
self.ticks.append(event)
for field in self.fields:
self.totals[field] += event[field]
# We return a list of all out-of-range events we removed.
out_of_range = []
# Clear out expired events, decrementing totals.
# newest oldest
# | |
# V V
while (self.ticks[-1].dt - self.ticks[0].dt) >= self.delta:
# popleft removes and returns ticks[0]
popped = self.ticks.popleft()
# Decrement totals
for field in self.fields:
self.totals[field] -= popped[field]
# Add the popped element to the list of dropped events.
out_of_range.append(popped)
return out_of_range
def average(self, field):
assert field in self.fields
if len(self.ticks) == 0:
return 0.0
else:
return self.totals[field] / len(self.ticks)
def get_averages(self):
"""
Return an ndict of all our tracked averages.
"""
out = ndict()
for field in self.fields:
out[field] = self.average(field)
return out
def assert_well_formed(self, event):
assert isinstance(event, ndict), "Bad event in EventWindow:%s" % event
assert event.has_key('dt'), "Missing dt in EventWindow:%s" % event
assert isinstance(event.dt, datetime),"Bad dt in EventWindow:%s" % event
if len(self.ticks) > 0:
# Something is wrong if new event is older than previous.
assert event.dt >= self.ticks[-1].dt, \
"Events arrived out of order in EventWindow: %s -> %s" % (event, self.ticks[0])
for field in self.fields:
assert event.has_key(field), \
"Event missing [%s] in EventWindow" % field
assert isinstance(event[field], Number), \
"Got %s for %s in EventWindow" % (event[field], field)
if __name__ == "__main__":
def make_event(**kwargs):
e = ndict()
for key, value in kwargs.iteritems():
e[key] = value
return e
dates = date_gen(delta = timedelta(hours = 12))
events = (
make_event(
sid = 'foo', price = random.random(),
dt = date,
type = zp.DATASOURCE_TYPE.TRADE,
source_id = 'ds',
vol = i
)
for date, i in izip(dates, xrange(100))
)
gen = MovingAverageTransformGen(events, 1, ['price', 'vol'])
+25 -3
View File
@@ -3,7 +3,7 @@ import numbers
from hashlib import md5
from datetime import datetime, timedelta
from itertools import izip_longest
from zipline import ndict
from zipline.protocol import DATASOURCE_TYPE
@@ -19,7 +19,14 @@ def mock_raw_event(sid, dt):
def date_gen(start = datetime(2012, 6, 6, 0), delta = timedelta(minutes = 1), n = 100):
return (start + i * delta for i in xrange(n))
def stringify_args(*args, **kwargs):
def alternate(g1, g2):
for e1, e2 in izip_longest(g1, g2):
if e1 != None:
yield e1
if e2 != None:
yield e2
def hash_args(*args, **kwargs):
"""Define a unique string for any set of representable args."""
arg_string = '_'.join([str(arg) for arg in args])
kwarg_string = '_'.join([str(key) + '=' + str(value) for key, value in kwargs.iteritems()])
@@ -55,11 +62,26 @@ def assert_trade_protocol(event):
def assert_datasource_unframe_protocol(event):
"""Assert that an event is valid output of zp.DATASOURCE_UNFRAME."""
assert isinstance(event, ndict)
assert isinstance(event.source_id, basestring)
assert event.type in DATASOURCE_TYPE
assert event.has_key('dt')
def assert_feed_protocol(event):
"""Assert that an event is valid input to zp.FEED_FRAME."""
assert isinstance(event, ndict)
assert isinstance(event.source_id, basestring)
assert event.type in DATASOURCE_TYPE
assert event.has_key('dt')
def assert_feed_unframe_protocol(event):
"""Same as above."""
assert isinstance(event, ndict)
assert event.type in DATASOURCE_TYPE
assert event.has_key('dt')
def assert_transform_protocol(event):
pass
+4
View File
@@ -318,6 +318,7 @@ def FEED_FRAME(event):
- source_id
- type
- dt
"""
assert isinstance(event, ndict), 'unknown type %s' % str(event)
source_id = event.source_id
@@ -332,6 +333,9 @@ def FEED_UNFRAME(msg):
#TODO: anything we can do to assert more about the content of the dict?
assert isinstance(payload, dict)
rval = ndict(payload)
assert rval.source_id
assert rval.type in DATASOURCE_TYPE
assert rval.dt
UNPACK_DATE(rval)
return rval
except TypeError: