mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-26 00:45:44 +08:00
Moved some code around between node_context_windows.py and context_windows.py
This commit is contained in:
parent
b38d13d9c8
commit
3a8be588a6
@ -7,8 +7,10 @@ from dataclasses import dataclass
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import logging
|
import logging
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
import comfy.patcher_extension
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_base import BaseModel
|
from comfy.model_base import BaseModel
|
||||||
|
from comfy.model_patcher import ModelPatcher
|
||||||
from comfy.controlnet import ControlBase
|
from comfy.controlnet import ControlBase
|
||||||
|
|
||||||
|
|
||||||
@ -249,6 +251,27 @@ class IndexListContextHandler(ContextHandlerABC):
|
|||||||
|
|
||||||
# TODO: add callback here
|
# TODO: add callback here
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_sampling_wrapper(executor, model, noise_shape: torch.Tensor, *args, **kwargs):
|
||||||
|
# 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: IndexListContextHandler = model_options.get("context_handler", None)
|
||||||
|
if handler is not None:
|
||||||
|
noise_shape = list(noise_shape)
|
||||||
|
noise_shape[handler.dim] = min(noise_shape[handler.dim], handler.context_length)
|
||||||
|
return executor(model, noise_shape, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def create_prepare_sampling_wrapper(model: ModelPatcher):
|
||||||
|
model.add_wrapper_with_key(
|
||||||
|
comfy.patcher_extension.WrappersMP.PREPARE_SAMPLING,
|
||||||
|
"ContextWindows_prepare_sampling",
|
||||||
|
_prepare_sampling_wrapper
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def match_weights_to_dim(weights: list[float], x_in: torch.Tensor, dim: int, device=None) -> torch.Tensor:
|
def match_weights_to_dim(weights: list[float], x_in: torch.Tensor, dim: int, device=None) -> torch.Tensor:
|
||||||
total_dims = len(x_in.shape)
|
total_dims = len(x_in.shape)
|
||||||
weights_tensor = torch.Tensor(weights).to(device=device)
|
weights_tensor = torch.Tensor(weights).to(device=device)
|
||||||
|
|||||||
@ -1,62 +1,10 @@
|
|||||||
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 comfy.samplers
|
|
||||||
import nodes
|
import nodes
|
||||||
import torch
|
|
||||||
|
|
||||||
|
|
||||||
def _prepare_sampling_wrapper(executor, model, noise_shape: torch.Tensor, *args, **kwargs):
|
class ContextWindowsManualNode(io.ComfyNode):
|
||||||
# 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 = list(noise_shape)
|
|
||||||
noise_shape[handler.dim] = min(noise_shape[handler.dim], handler.context_length)
|
|
||||||
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)
|
|
||||||
|
|
||||||
def _outer_sample_wrapper(executor, *args, **kwargs):
|
|
||||||
guider: comfy.samplers.CFGGuider = executor.class_obj
|
|
||||||
handler: comfy.context_windows.IndexListContextHandler = guider.model_options.get("context_handler", None)
|
|
||||||
if handler is not None:
|
|
||||||
args = list(args)
|
|
||||||
noise: torch.Tensor = args[0]
|
|
||||||
length = noise.shape[handler.dim]
|
|
||||||
window = comfy.context_windows.IndexListContextWindow(list(range(handler.context_length)))
|
|
||||||
noise = window.get_tensor(noise, dim=handler.dim)
|
|
||||||
cat_count = (length // handler.context_length) + 1
|
|
||||||
noise = torch.cat([noise] * cat_count, dim=handler.dim)
|
|
||||||
if handler.dim == 0:
|
|
||||||
noise = noise[:length]
|
|
||||||
elif handler.dim == 1:
|
|
||||||
noise = noise[:, :length]
|
|
||||||
elif handler.dim == 2:
|
|
||||||
noise = noise[:, :, :length]
|
|
||||||
else:
|
|
||||||
pass
|
|
||||||
args[0] = noise
|
|
||||||
args = tuple(args)
|
|
||||||
return executor(*args, **kwargs)
|
|
||||||
|
|
||||||
def create_outer_sampler_wrapper(model_options: dict):
|
|
||||||
comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
|
||||||
"ContextWindows_outer_sample",
|
|
||||||
_outer_sample_wrapper,
|
|
||||||
model_options, is_model_options=True)
|
|
||||||
|
|
||||||
|
|
||||||
class ContextWindowsNode(io.ComfyNode):
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def define_schema(cls) -> io.Schema:
|
def define_schema(cls) -> io.Schema:
|
||||||
return io.Schema(
|
return io.Schema(
|
||||||
@ -96,12 +44,11 @@ class ContextWindowsNode(io.ComfyNode):
|
|||||||
context_stride=context_stride,
|
context_stride=context_stride,
|
||||||
closed_loop=closed_loop,
|
closed_loop=closed_loop,
|
||||||
dim=dim)
|
dim=dim)
|
||||||
create_prepare_sampling_wrapper(model.model_options)
|
# make memory usage calculation only take into account the context window latents
|
||||||
#create_outer_sampler_wrapper(model.model_options)
|
comfy.context_windows.create_prepare_sampling_wrapper(model)
|
||||||
return io.NodeOutput(model)
|
return io.NodeOutput(model)
|
||||||
|
|
||||||
|
class WanContextWindowsNode(ContextWindowsManualNode):
|
||||||
class WanContextWindowsNode(ContextWindowsNode):
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def define_schema(cls) -> io.Schema:
|
def define_schema(cls) -> io.Schema:
|
||||||
schema = super().define_schema()
|
schema = super().define_schema()
|
||||||
@ -127,10 +74,11 @@ class WanContextWindowsNode(ContextWindowsNode):
|
|||||||
context_overlap = max(((context_overlap - 1) // 4) + 1, 0) # at least overlap 0
|
context_overlap = max(((context_overlap - 1) // 4) + 1, 0) # at least overlap 0
|
||||||
return super().execute(model, context_length, context_overlap, context_schedule, context_stride, closed_loop, fuse_method, 2)
|
return super().execute(model, context_length, context_overlap, context_schedule, context_stride, closed_loop, fuse_method, 2)
|
||||||
|
|
||||||
|
|
||||||
class ContextWindowsExtension(ComfyExtension):
|
class ContextWindowsExtension(ComfyExtension):
|
||||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||||
return [
|
return [
|
||||||
ContextWindowsNode,
|
ContextWindowsManualNode,
|
||||||
WanContextWindowsNode,
|
WanContextWindowsNode,
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user