mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-24 02:31:18 +08:00
Add prepare_sampling wrapper for context window to more accurately estimate latent memory requirements, fixed merging wrappers/callbacks dicts in prepare_model_patcher
This commit is contained in:
parent
c5c5824daf
commit
58152b5606
@ -149,7 +149,7 @@ def cleanup_models(conds, models):
|
|||||||
|
|
||||||
cleanup_additional_models(set(control_cleanup))
|
cleanup_additional_models(set(control_cleanup))
|
||||||
|
|
||||||
def prepare_model_patcher(model: 'ModelPatcher', conds, model_options: dict):
|
def prepare_model_patcher(model: ModelPatcher, conds, model_options: dict):
|
||||||
'''
|
'''
|
||||||
Registers hooks from conds.
|
Registers hooks from conds.
|
||||||
'''
|
'''
|
||||||
@ -158,8 +158,8 @@ def prepare_model_patcher(model: 'ModelPatcher', conds, model_options: dict):
|
|||||||
for k in conds:
|
for k in conds:
|
||||||
get_hooks_from_cond(conds[k], hooks)
|
get_hooks_from_cond(conds[k], hooks)
|
||||||
# add wrappers and callbacks from ModelPatcher to transformer_options
|
# add wrappers and callbacks from ModelPatcher to transformer_options
|
||||||
model_options["transformer_options"]["wrappers"] = comfy.patcher_extension.copy_nested_dicts(model.wrappers)
|
comfy.patcher_extension.merge_nested_dicts(model_options["transformer_options"].setdefault("wrappers", {}), model.wrappers)
|
||||||
model_options["transformer_options"]["callbacks"] = comfy.patcher_extension.copy_nested_dicts(model.callbacks)
|
comfy.patcher_extension.merge_nested_dicts(model_options["transformer_options"].setdefault("callbacks", {}), model.callbacks)
|
||||||
# begin registering hooks
|
# begin registering hooks
|
||||||
registered = comfy.hooks.HookGroup()
|
registered = comfy.hooks.HookGroup()
|
||||||
target_dict = comfy.hooks.create_target_dict(comfy.hooks.EnumWeightTarget.Model)
|
target_dict = comfy.hooks.create_target_dict(comfy.hooks.EnumWeightTarget.Model)
|
||||||
|
|||||||
@ -1,6 +1,27 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
from comfy_api.latest import ComfyExtension, io
|
from comfy_api.latest import ComfyExtension, io
|
||||||
import comfy.context_windows
|
import comfy.context_windows
|
||||||
|
import comfy.patcher_extension
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_sampling_wrapper(executor, model, noise_shape: torch.Tensor, *args, **kwargs):
|
||||||
|
# TODO: handle various dims instead of defaulting to 0th
|
||||||
|
# limit noise_shape length to context_length for more accurate vram use estimation
|
||||||
|
model_options = kwargs.get("model_options", None)
|
||||||
|
if model_options is None:
|
||||||
|
raise Exception("model_options not found in prepare_sampling_wrapper; this should never happen, something went wrong.")
|
||||||
|
handler: comfy.context_windows.IndexListContextHandler = model_options.get("context_handler", None)
|
||||||
|
if handler is not None:
|
||||||
|
noise_shape = [min(noise_shape[0], handler.context_length)] + list(noise_shape[1:])
|
||||||
|
return executor(model, noise_shape, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def create_prepare_sampling_wrapper(model_options: dict):
|
||||||
|
comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.PREPARE_SAMPLING,
|
||||||
|
"ContextWindows_prepare_sampling",
|
||||||
|
_prepare_sampling_wrapper,
|
||||||
|
model_options, is_model_options=True)
|
||||||
|
|
||||||
|
|
||||||
class ContextWindowsNode(io.ComfyNode):
|
class ContextWindowsNode(io.ComfyNode):
|
||||||
@ -37,6 +58,7 @@ class ContextWindowsNode(io.ComfyNode):
|
|||||||
context_length=context_length,
|
context_length=context_length,
|
||||||
context_overlap=context_overlap,
|
context_overlap=context_overlap,
|
||||||
dim=dim)
|
dim=dim)
|
||||||
|
create_prepare_sampling_wrapper(model.model_options)
|
||||||
return io.NodeOutput(model)
|
return io.NodeOutput(model)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user