From c399dec25785c8f9eab402df67af38d74c746084 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 3 Jan 2024 18:07:50 -0800 Subject: [PATCH] In .step() expose only values of lastvalue channels --- permchain/channels/base.py | 11 ----------- permchain/pregel/__init__.py | 17 ++++++++++++++--- 2 files changed, 14 insertions(+), 14 deletions(-) diff --git a/permchain/channels/base.py b/permchain/channels/base.py index ab951e42e..643f77f87 100644 --- a/permchain/channels/base.py +++ b/permchain/channels/base.py @@ -129,14 +129,3 @@ def create_checkpoint( except EmptyChannelError: pass return checkpoint - - -def channel_values(channels: Mapping[str, BaseChannel]) -> dict[str, Any]: - """Return a dictionary of channel values.""" - values: dict[str, Any] = {} - for k, v in channels.items(): - try: - values[k] = v.get() - except EmptyChannelError: - pass - return values diff --git a/permchain/pregel/__init__.py b/permchain/pregel/__init__.py index eff58980b..1eee4933b 100644 --- a/permchain/pregel/__init__.py +++ b/permchain/pregel/__init__.py @@ -45,7 +45,6 @@ from permchain.channels.base import ( BaseChannel, ChannelsManager, EmptyChannelError, - channel_values, create_checkpoint, ) from permchain.channels.last_value import LastValue @@ -281,7 +280,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): # yield current value and checkpoint view view = CheckpointView( - values=channel_values(channels), + values=_updateable_channel_values(channels), step=step + 1, ) yield map_output(self.output, pending_writes, channels), view @@ -379,7 +378,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): # yield current value and checkpoint view view = CheckpointView( - values=channel_values(channels), + values=_updateable_channel_values(channels), step=step + 1, ) yield map_output(self.output, pending_writes, channels), view @@ -628,3 +627,15 @@ def _prepare_next_tasks( seen[proc.channel] = checkpoint["channel_versions"][proc.channel] return tasks + + +def _updateable_channel_values(channels: Mapping[str, BaseChannel]) -> dict[str, Any]: + """Return a dictionary of updateable channel values.""" + values: dict[str, Any] = {} + for k, v in channels.items(): + if isinstance(v, LastValue): + try: + values[k] = v.get() + except EmptyChannelError: + pass + return values