Add checkpointer=True mode for subgraphs that want to keep state between turns (#3055)

This commit is contained in:
Nuno Campos
2025-01-15 16:08:03 -08:00
committed by GitHub
3 changed files with 93 additions and 9 deletions
+28 -6
View File
@@ -494,7 +494,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# get the subgraphs
@@ -606,7 +608,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# get the subgraphs
@@ -926,7 +930,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# apply null writes
@@ -1020,7 +1026,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# apply null writes
@@ -1209,7 +1217,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# apply null writes
@@ -1303,7 +1313,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1,
for_execution=True,
store=self.store,
checkpointer=self.checkpointer or None,
checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None,
)
# apply null writes
@@ -1455,6 +1467,8 @@ class Pregel(PregelProtocol):
checkpointer: Optional[BaseCheckpointSaver] = None
elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}):
checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER]
elif self.checkpointer is True:
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
else:
checkpointer = self.checkpointer
if checkpointer and not config.get(CONF):
@@ -1598,6 +1612,12 @@ class Pregel(PregelProtocol):
interrupt_after=interrupt_after,
debug=debug,
)
# set up subgraph checkpointing
if self.checkpointer is True:
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = NS_SEP.join(
part.split(NS_END)[0] for part in ns.split(NS_SEP)
)
# set up messages stream mode
if "messages" in stream_modes:
run_manager.inheritable_handlers.append(
@@ -1622,6 +1642,7 @@ class Pregel(PregelProtocol):
interrupt_after=interrupt_after_,
manager=run_manager,
debug=debug,
check_subgraphs=self.checkpointer is not True,
) as loop:
# create runner
runner = PregelRunner(
@@ -1849,6 +1870,7 @@ class Pregel(PregelProtocol):
interrupt_after=interrupt_after_,
manager=run_manager,
debug=debug,
check_subgraphs=self.checkpointer is not True,
) as loop:
# create runner
runner = PregelRunner(
+5 -3
View File
@@ -38,9 +38,11 @@ except ImportError:
All = Literal["*"]
"""Special value to indicate that graph should interrupt on all nodes."""
Checkpointer = Union[None, Literal[False], BaseCheckpointSaver]
"""Type of the checkpointer to use for a subgraph. False disables checkpointing,
even if the parent graph has a checkpointer. None inherits checkpointer."""
Checkpointer = Union[None, bool, BaseCheckpointSaver]
"""Type of the checkpointer to use for a subgraph.
- True enables persistent checkpointing for this subgraph.
- False disables checkpointing, even if the parent graph has a checkpointer.
- None inherits checkpointer from the parent graph."""
StreamMode = Literal["values", "updates", "debug", "messages", "custom"]
"""How the stream method should emit outputs.
+60
View File
@@ -3189,6 +3189,66 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_subgraph_checkpoint_true(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name)
class InnerState(TypedDict):
my_key: Annotated[str, operator.add]
my_other_key: str
def inner_1(state: InnerState):
return {"my_key": " got here", "my_other_key": state["my_key"]}
def inner_2(state: InnerState):
return {"my_key": " and there"}
inner = StateGraph(InnerState)
inner.add_node("inner_1", inner_1)
inner.add_node("inner_2", inner_2)
inner.add_edge("inner_1", "inner_2")
inner.set_entry_point("inner_1")
inner.set_finish_point("inner_2")
class State(TypedDict):
my_key: str
graph = StateGraph(State)
graph.add_node("inner", inner.compile(checkpointer=True))
graph.add_edge(START, "inner")
graph.add_conditional_edges(
"inner", lambda s: "inner" if s["my_key"].count("there") < 2 else END
)
app = graph.compile(checkpointer=checkpointer)
config = {"configurable": {"thread_id": "2"}}
assert [c for c in app.stream({"my_key": ""}, config, subgraphs=True)] == [
(("inner",), {"inner_1": {"my_key": " got here", "my_other_key": ""}}),
(("inner",), {"inner_2": {"my_key": " and there"}}),
((), {"inner": {"my_key": " got here and there"}}),
(
("inner",),
{
"inner_1": {
"my_key": " got here",
"my_other_key": " got here and there got here and there",
}
},
),
(("inner",), {"inner_2": {"my_key": " and there"}}),
(
(),
{
"inner": {
"my_key": " got here and there got here and there got here and there"
}
},
),
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_stream_subgraphs_during_execution(
request: pytest.FixtureRequest, checkpointer_name: str