optimize subgraph state lookups

This commit is contained in:
vbarda
2024-08-22 20:35:08 -04:00
parent 9f6e57d2a7
commit 4162be8119
4 changed files with 116 additions and 115 deletions
@@ -255,6 +255,9 @@ class BasePostgresSaver(BaseCheckpointSaver):
if config:
wheres.append("thread_id = %s ")
param_values.append(config["configurable"]["thread_id"])
if checkpoint_id := get_checkpoint_id(config):
wheres.append("checkpoint_id = %s ")
param_values.append(checkpoint_id)
# construct predicate for metadata filter
if filter:
@@ -70,6 +70,9 @@ def search_where(
if config is not None:
wheres.append("thread_id = ?")
param_values.append(config["configurable"]["thread_id"])
if checkpoint_id := get_checkpoint_id(config):
wheres.append("checkpoint_id = ?")
param_values.append(checkpoint_id)
# construct predicate for metadata filter
if filter:
@@ -177,6 +177,7 @@ class MemorySaver(
Iterator[CheckpointTuple]: An iterator of matching checkpoint tuples.
"""
thread_ids = (config["configurable"]["thread_id"],) if config else self.storage
config_checkpoint_id = get_checkpoint_id(config) if config else None
for thread_id in thread_ids:
for checkpoint_ns in self.storage[thread_id].keys():
for checkpoint_id, (
@@ -188,7 +189,11 @@ class MemorySaver(
key=lambda x: x[0],
reverse=True,
):
# filter by checkpoint ID
# filter by checkpoint ID from config
if config_checkpoint_id and checkpoint_id != config_checkpoint_id:
continue
# filter by checkpoint ID from `before` config
if (
before
and (before_checkpoint_id := get_checkpoint_id(before))
+104 -114
View File
@@ -178,34 +178,31 @@ class Channel:
)
def _get_subgraph(graph: Pregel, checkpoint_ns: str) -> Pregel:
if checkpoint_ns == "":
return graph
path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
nodes = graph.nodes
for subgraph_node_name in path:
# if we have this separator it means we have a node that was triggered by Send
if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name:
name_parts = subgraph_node_name.split(SEND_CHECKPOINT_NAMESPACE_SEPARATOR)
if len(name_parts) != 2:
raise ValueError(f"Malformed node name '{subgraph_node_name}'")
subgraph_node_name = name_parts[0]
if subgraph_node_name not in nodes:
raise ValueError(f"Couldn't find node '{subgraph_node_name}'.")
subgraph_node = nodes[subgraph_node_name]
if isinstance(subgraph_node.bound, Pregel):
nodes = subgraph_node.bound.nodes
elif isinstance(subgraph_node.bound, RunnableSequence):
for runnable in subgraph_node.bound.steps:
def _get_checkpoint_ns_to_graph(
graph: Pregel, checkpoint_ns_to_graph: dict[str, Pregel] = {}, checkpoint_ns=""
) -> Pregel:
for node_name, node in graph.nodes.items():
if isinstance(node.bound, Pregel):
_get_checkpoint_ns_to_graph(
node.bound,
checkpoint_ns_to_graph,
f"{checkpoint_ns}{CHECKPOINT_NAMESPACE_SEPARATOR}{node_name}"
if checkpoint_ns
else node_name,
)
elif isinstance(node.bound, RunnableSequence):
for runnable in node.bound.steps:
if isinstance(runnable, Pregel):
nodes = runnable.nodes
break
else:
continue
return subgraph_node.bound
_get_checkpoint_ns_to_graph(
node.bound,
checkpoint_ns_to_graph,
f"{checkpoint_ns}{CHECKPOINT_NAMESPACE_SEPARATOR}{node_name}"
if checkpoint_ns
else node_name,
)
checkpoint_ns_to_graph[checkpoint_ns] = graph
return checkpoint_ns_to_graph
def _has_nested_interrupts(
@@ -401,59 +398,54 @@ class Pregel(
saved = self.checkpointer.get_tuple(config)
checkpoint_config = saved.config if saved else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
for saved in self.checkpointer.list(config):
checkpoint_ns_to_graph: dict[str, Pregel] = _get_checkpoint_ns_to_graph(self)
# we only lookup subgraph checkpoints if we actually have subgraphs
if len(set(checkpoint_ns_to_graph)) == 1:
checkpoint_tuples = (saved,)
else:
checkpoint_tuples = self.checkpointer.list(saved.config)
for saved in checkpoint_tuples:
saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"]
saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"]
if checkpoint_id != saved_checkpoint_id:
graph_checkpoint_ns = saved_checkpoint_ns.split(
SEND_CHECKPOINT_NAMESPACE_SEPARATOR
)[0]
graph = checkpoint_ns_to_graph.get(graph_checkpoint_ns)
if graph is None:
continue
existing_checkpoint_id = checkpoint_ns_to_checkpoint_id.get(
saved_checkpoint_ns
)
# keep only most recent checkpoint_id
if (
existing_checkpoint_id is None
or saved_checkpoint_id > existing_checkpoint_id
with ChannelsManager(
graph.channels, saved.checkpoint, saved.config, skip_context=True
) as (
channels,
managed,
):
if saved_checkpoint_ns not in checkpoint_ns_to_graph:
checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph(
self, saved_checkpoint_ns
)
graph = checkpoint_ns_to_graph[saved_checkpoint_ns]
with ChannelsManager(
graph.channels, saved.checkpoint, saved.config, skip_context=True
) as (
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
):
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
saved_checkpoint_ns
] = saved_checkpoint_id
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[saved_checkpoint_ns] = saved_checkpoint_id
if not checkpoint_ns_to_state_snapshots:
return StateSnapshot(
@@ -479,57 +471,55 @@ class Pregel(
saved = await self.checkpointer.aget_tuple(config)
checkpoint_config = saved.config if saved else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
async for saved in self.checkpointer.alist(config):
checkpoint_ns_to_graph: dict[str, Pregel] = _get_checkpoint_ns_to_graph(self)
# we only lookup subgraph checkpoints if we actually have subgraphs
if len(set(checkpoint_ns_to_graph)) == 1:
async def list_checkpoints():
yield saved
checkpoint_tuples = list_checkpoints()
else:
checkpoint_tuples = self.checkpointer.alist(saved.config)
async for saved in checkpoint_tuples:
saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"]
saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"]
if checkpoint_id != saved_checkpoint_id:
graph_checkpoint_ns = saved_checkpoint_ns.split(
SEND_CHECKPOINT_NAMESPACE_SEPARATOR
)[0]
graph = checkpoint_ns_to_graph.get(graph_checkpoint_ns)
if graph is None:
continue
existing_checkpoint_id = checkpoint_ns_to_checkpoint_id.get(
saved_checkpoint_ns
)
async with AsyncChannelsManager(
graph.channels, saved.checkpoint, saved.config, skip_context=True
) as (channels, managed):
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
# keep only most recent checkpoint_id
if (
existing_checkpoint_id is None
or saved_checkpoint_id > existing_checkpoint_id
):
if saved_checkpoint_ns not in checkpoint_ns_to_graph:
checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph(
self, saved_checkpoint_ns
)
graph = checkpoint_ns_to_graph[saved_checkpoint_ns]
async with AsyncChannelsManager(
graph.channels, saved.checkpoint, saved.config, skip_context=True
) as (channels, managed):
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
saved_checkpoint_ns
] = saved_checkpoint_id
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[saved_checkpoint_ns] = saved_checkpoint_id
if not checkpoint_ns_to_state_snapshots:
return StateSnapshot(