In .step() expose only values of lastvalue channels

This commit is contained in:
Nuno Campos
2024-01-03 18:07:50 -08:00
parent b14c2638ee
commit c399dec257
2 changed files with 14 additions and 14 deletions
-11
View File
@@ -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
+14 -3
View File
@@ -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