From 7051d6826dc8b4f634d9a86a0b77b05de57deed1 Mon Sep 17 00:00:00 2001 From: scottsanderson Date: Tue, 21 Aug 2012 12:18:26 -0400 Subject: [PATCH] remove unnecessary copying --- zipline/gens/composites.py | 2 +- zipline/gens/transform.py | 25 ++++++++++++++++--------- 2 files changed, 17 insertions(+), 10 deletions(-) diff --git a/zipline/gens/composites.py b/zipline/gens/composites.py index b3fa7576..78bea152 100644 --- a/zipline/gens/composites.py +++ b/zipline/gens/composites.py @@ -43,7 +43,7 @@ def merged_transforms(sorted_stream, *transforms): """ for transform in transforms: assert isinstance(transform, StatefulTransform) - + transform.set_copying() # Generate expected hashes for each transform namestrings = [tnfm.get_hash() for tnfm in transforms] diff --git a/zipline/gens/transform.py b/zipline/gens/transform.py index 51475f45..612cd098 100644 --- a/zipline/gens/transform.py +++ b/zipline/gens/transform.py @@ -72,6 +72,7 @@ class StatefulTransform(object): # Create an instance of our transform class. self.state = tnfm_class(*args, **kwargs) + self._copying = False # Create the string associated with this generator's output. self.namestring = tnfm_class.__name__ + hash_args(*args, **kwargs) @@ -80,12 +81,16 @@ class StatefulTransform(object): def get_hash(self): return self.namestring + def set_copyting(self): + self._copying = True + def transform(self, stream_in): return self._gen(stream_in) def _gen(self, stream_in): # IMPORTANT: Messages may contain pointers that are shared with - # other streams, so we only manipulate copies. + # other streams. Transforms that modify their input + # messages should only manipulate copies. log.info('Running StatefulTransform [%s]' % self.get_hash()) for message in stream_in: @@ -94,13 +99,15 @@ class StatefulTransform(object): if message == None: continue - #TODO: refactor this to avoid unnecessary copying. - assert_sort_unframe_protocol(message) - message_copy = deepcopy(message) - + + # Copying flag is used by merged_transforms to ensure + # isolation of messages. + if self._copying: + message = deepcopy(message) + # Same shared pointer issue here as above. - tnfm_value = self.state.update(deepcopy(message_copy)) + tnfm_value = self.state.update(message) # FORWARDER flag means we want to keep all original # values, plus append tnfm_id and tnfm_value. Used for @@ -108,7 +115,7 @@ class StatefulTransform(object): # will be fed into a merge. Currently only Passthrough # uses this flag. if self.forward_all: - out_message = message_copy + out_message = message out_message.tnfm_id = self.namestring out_message.tnfm_value = tnfm_value yield out_message @@ -133,7 +140,7 @@ class StatefulTransform(object): # tnfm_value) into a single field. This mode is used by # the sequential_transforms composite. elif self.append_value: - out_message = message_copy + out_message = message out_message[self.namestring] = tnfm_value yield out_message @@ -145,7 +152,7 @@ class StatefulTransform(object): out_message = ndict() out_message.tnfm_id = self.namestring out_message.tnfm_value = tnfm_value - out_message.dt = message_copy.dt + out_message.dt = message.dt yield out_message log.info('Finished StatefulTransform [%s]' % self.get_hash())