Skip to content

Commit bb2c530

Browse files
authored
Optimize iterator graph materialization (#9355)
* Optimize iterator graph materialization * Update generated OpenAPI schema * Optimize iterator graph expansion and memory * Fix graph adjacency cache lifecycle * Refactored `try...finally` behavior * Optimize saved workflow graph restoration
1 parent 1aeb05b commit bb2c530

16 files changed

Lines changed: 1375 additions & 156 deletions

invokeai/app/invocations/baseinvocation.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -207,6 +207,10 @@ def invoke(self, context: InvocationContext) -> BaseInvocationOutput:
207207
"""Invoke with provided context and return outputs."""
208208
pass
209209

210+
def get_event_invocation(self) -> "BaseInvocation":
211+
"""Returns the invocation representation included in execution events."""
212+
return self
213+
210214
def invoke_internal(self, context: InvocationContext, services: "InvocationServices") -> BaseInvocationOutput:
211215
"""
212216
Internal invoke method, calls `invoke()` after some prep.

invokeai/app/services/events/events_base.py

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -64,7 +64,7 @@ def dispatch(self, event: "EventBase") -> None:
6464

6565
def emit_invocation_started(self, queue_item: "SessionQueueItem", invocation: "BaseInvocation") -> None:
6666
"""Emitted when an invocation is started"""
67-
self.dispatch(InvocationStartedEvent.build(queue_item, invocation))
67+
self.dispatch(InvocationStartedEvent.build(queue_item, invocation.get_event_invocation()))
6868

6969
def emit_invocation_progress(
7070
self,
@@ -75,13 +75,15 @@ def emit_invocation_progress(
7575
image: "ProgressImage | None" = None,
7676
) -> None:
7777
"""Emitted at periodically during an invocation"""
78-
self.dispatch(InvocationProgressEvent.build(queue_item, invocation, message, percentage, image))
78+
self.dispatch(
79+
InvocationProgressEvent.build(queue_item, invocation.get_event_invocation(), message, percentage, image)
80+
)
7981

8082
def emit_invocation_complete(
8183
self, queue_item: "SessionQueueItem", invocation: "BaseInvocation", output: "BaseInvocationOutput"
8284
) -> None:
8385
"""Emitted when an invocation is complete"""
84-
self.dispatch(InvocationCompleteEvent.build(queue_item, invocation, output))
86+
self.dispatch(InvocationCompleteEvent.build(queue_item, invocation.get_event_invocation(), output))
8587

8688
def emit_invocation_error(
8789
self,
@@ -92,7 +94,15 @@ def emit_invocation_error(
9294
error_traceback: str,
9395
) -> None:
9496
"""Emitted when an invocation encounters an error"""
95-
self.dispatch(InvocationErrorEvent.build(queue_item, invocation, error_type, error_message, error_traceback))
97+
self.dispatch(
98+
InvocationErrorEvent.build(
99+
queue_item,
100+
invocation.get_event_invocation(),
101+
error_type,
102+
error_message,
103+
error_traceback,
104+
)
105+
)
96106

97107
# endregion
98108

invokeai/app/services/session_processor/session_processor_default.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@
3333
WorkflowCallQueueLifecycle,
3434
)
3535
from invokeai.app.services.session_queue.session_queue_common import SessionQueueItem, SessionQueueItemNotFoundError
36-
from invokeai.app.services.shared.graph import NodeInputError
36+
from invokeai.app.services.shared.graph import CollectInvocation, IterateInvocation, NodeInputError
3737
from invokeai.app.services.shared.invocation_context import InvocationContextData, build_invocation_context
3838
from invokeai.app.util.profiler import Profiler
3939

@@ -141,10 +141,19 @@ def run_node(self, invocation: BaseInvocation, queue_item: SessionQueueItem):
141141

142142
# Invoke the node
143143
output = invocation.invoke_internal(context=context, services=self._services)
144+
control_collection = None
145+
if self._on_after_run_node_callbacks and isinstance(invocation, (IterateInvocation, CollectInvocation)):
146+
control_collection = invocation.collection
144147
# Save output and history
145148
queue_item.session.complete(invocation.id, output)
146149

147-
self._on_after_run_node(invocation, queue_item, output)
150+
if control_collection is not None:
151+
invocation.collection = control_collection
152+
try:
153+
self._on_after_run_node(invocation, queue_item, output)
154+
finally:
155+
if control_collection is not None:
156+
invocation.collection = []
148157

149158
except CanceledException:
150159
# A CanceledException is raised during the denoising step callback if the cancel event is set. We don't need
@@ -215,7 +224,7 @@ def _on_after_run_session(self, queue_item: SessionQueueItem) -> None:
215224
# The queue item may have been canceled or failed while the session was running. We should only complete it
216225
# if it is not already canceled or failed.
217226
if queue_item.status not in ["canceled", "failed"] and queue_item.session.is_complete():
218-
queue_item = self._services.session_queue.complete_queue_item(queue_item.item_id)
227+
queue_item = self._services.session_queue.complete_queue_item(queue_item.item_id, queue_item=queue_item)
219228

220229
# We'll get a GESStatsNotFoundError if we try to log stats for an untracked graph, but in the processor
221230
# we don't care about that - suppress the error.

invokeai/app/services/session_processor/workflow_call_runtime.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -96,7 +96,7 @@ def begin_workflow_call_boundary(
9696
child_queue_item = None
9797
enqueued_child_item_ids: list[int] = []
9898
try:
99-
self._session_runner._services.session_queue.set_queue_item_session(queue_item.item_id, queue_item.session)
99+
self._session_runner._services.session_queue.save_queue_item_session(queue_item.item_id, queue_item.session)
100100
for child_result in child_session_results:
101101
child_queue_item = self._session_runner._services.session_queue.enqueue_workflow_call_child(
102102
parent_queue_item=queue_item,
@@ -105,8 +105,8 @@ def begin_workflow_call_boundary(
105105
)
106106
enqueued_child_item_ids.append(child_queue_item.item_id)
107107
queue_item.session.set_waiting_workflow_call_child_item_ids(enqueued_child_item_ids)
108-
self._session_runner._services.session_queue.set_queue_item_session(queue_item.item_id, queue_item.session)
109-
self._session_runner._services.session_queue.suspend_queue_item(queue_item.item_id)
108+
self._session_runner._services.session_queue.save_queue_item_session(queue_item.item_id, queue_item.session)
109+
self._session_runner._services.session_queue.suspend_queue_item(queue_item.item_id, queue_item=queue_item)
110110
except Exception as e:
111111
if enqueued_child_item_ids:
112112
self._session_runner._services.session_queue.delete_queue_items_by_id(enqueued_child_item_ids)
@@ -218,7 +218,7 @@ def _resume_parent_from_completed_child(self, child_queue_item: SessionQueueItem
218218
self._fail_parent_from_failed_child(parent_queue_item)
219219
return
220220
if not should_resume_parent:
221-
self._session_runner._services.session_queue.set_queue_item_session(
221+
self._session_runner._services.session_queue.save_queue_item_session(
222222
parent_queue_item.item_id, parent_queue_item.session
223223
)
224224
return
@@ -228,17 +228,19 @@ def _resume_parent_from_completed_child(self, child_queue_item: SessionQueueItem
228228
parent_output = WorkflowReturnOutput(values=aggregated_values)
229229
parent_queue_item.session.complete(waiting_invocation.id, parent_output)
230230
self._session_runner._on_after_run_node(waiting_invocation, parent_queue_item, parent_output)
231-
parent_queue_item = self._session_runner._services.session_queue.set_queue_item_session(
231+
self._session_runner._services.session_queue.save_queue_item_session(
232232
parent_queue_item.item_id, parent_queue_item.session
233233
)
234234
if parent_queue_item.session.is_complete():
235235
parent_queue_item = self._session_runner._services.session_queue.complete_queue_item(
236-
parent_queue_item.item_id
236+
parent_queue_item.item_id, queue_item=parent_queue_item
237237
)
238238
if getattr(parent_queue_item, "parent_item_id", None) is not None:
239239
self._resume_parent_from_completed_child(parent_queue_item)
240240
return
241-
self._session_runner._services.session_queue.resume_queue_item(parent_queue_item.item_id)
241+
self._session_runner._services.session_queue.resume_queue_item(
242+
parent_queue_item.item_id, queue_item=parent_queue_item
243+
)
242244

243245
def _fail_parent_from_failed_child(self, child_queue_item: SessionQueueItem) -> None:
244246
parent_queue_item = self._get_parent_queue_item(child_queue_item)

invokeai/app/services/session_queue/session_queue_base.py

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -92,8 +92,8 @@ def get_queue_status(
9292
acting_user_id is independent of user_id and controls only current-item redaction:
9393
when set, the returned status omits item_id/session_id/batch_id unless the
9494
currently-running item belongs to acting_user_id. The redaction is decided from the
95-
same get_current() snapshot used to embed those identifiers, so it cannot race against
96-
a concurrent state change.
95+
same database snapshot used to embed those identifiers, so it cannot race against a
96+
concurrent state change.
9797
9898
is_admin disables current-item redaction entirely: admins may see the identifiers of
9999
any user's current item. Redaction stays fail-closed - a caller that passes user_id
@@ -114,17 +114,17 @@ def get_batch_status(self, queue_id: str, batch_id: str, user_id: Optional[str]
114114
pass
115115

116116
@abstractmethod
117-
def complete_queue_item(self, item_id: int) -> SessionQueueItem:
117+
def complete_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
118118
"""Completes a session queue item"""
119119
pass
120120

121121
@abstractmethod
122-
def suspend_queue_item(self, item_id: int) -> SessionQueueItem:
122+
def suspend_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
123123
"""Suspends a session queue item while waiting on a child workflow execution."""
124124
pass
125125

126126
@abstractmethod
127-
def resume_queue_item(self, item_id: int) -> SessionQueueItem:
127+
def resume_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
128128
"""Resumes a suspended session queue item by returning it to pending state."""
129129
pass
130130

@@ -228,6 +228,11 @@ def set_queue_item_session(self, item_id: int, session: GraphExecutionState) ->
228228
"""Sets the session for a session queue item. Use this to update the session state."""
229229
pass
230230

231+
@abstractmethod
232+
def save_queue_item_session(self, item_id: int, session: GraphExecutionState) -> None:
233+
"""Persists a queue item's session without loading and returning the full queue item."""
234+
pass
235+
231236
@abstractmethod
232237
def enqueue_workflow_call_child(
233238
self,

0 commit comments

Comments
 (0)