Made context schedule and fuse method functions be stored on the handler instead of needing to be registered in core code to be found

This commit is contained in:
Jedrzej Kosinski 2025-08-07 16:31:11 -07:00
parent 1e3bd926f5
commit b38d13d9c8
2 changed files with 65 additions and 71 deletions

View File

@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Callable
import torch import torch
import numpy as np import numpy as np
import collections import collections
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
@ -76,17 +77,27 @@ class IndexListCallbacks:
return {} return {}
@dataclass
class ContextSchedule:
name: str
func: Callable
@dataclass
class ContextFuseMethod:
name: str
func: Callable
ContextResults = collections.namedtuple("ContextResults", ['window_idx', 'sub_conds_out', 'sub_conds', 'window']) ContextResults = collections.namedtuple("ContextResults", ['window_idx', 'sub_conds_out', 'sub_conds', 'window'])
class IndexListContextHandler(ContextHandlerABC): class IndexListContextHandler(ContextHandlerABC):
def __init__(self, context_schedule: str, fuse_method: str, context_length: int=1, context_overlap: int=0, dim=0): def __init__(self, context_schedule: ContextSchedule, fuse_method: ContextFuseMethod, context_length: int=1, context_overlap: int=0, context_stride: int=1, closed_loop=False, dim=0):
self.context_schedule = context_schedule self.context_schedule = context_schedule
self.fuse_method = fuse_method self.fuse_method = fuse_method
self.context_length = context_length self.context_length = context_length
self.context_overlap = context_overlap self.context_overlap = context_overlap
self.context_stride = 1 self.context_stride = context_stride
self.closed_loop = False self.closed_loop = closed_loop
self.step = 0 # TODO: get from model options
self.dim = dim self.dim = dim
self._step = 0
self.callbacks = {} self.callbacks = {}
@ -98,8 +109,6 @@ class IndexListContextHandler(ContextHandlerABC):
return False return False
def prepare_control_objects(self, control: ControlBase, device=None) -> ControlBase: def prepare_control_objects(self, control: ControlBase, device=None) -> ControlBase:
if device is not None:
control = control.get_instance_for_device(device)
if control.previous_controlnet is not None: if control.previous_controlnet is not None:
self.prepare_control_objects(control.previous_controlnet, device) self.prepare_control_objects(control.previous_controlnet, device)
return control return control
@ -149,30 +158,37 @@ class IndexListContextHandler(ContextHandlerABC):
resized_cond.append(resized_actual_cond) resized_cond.append(resized_actual_cond)
return resized_cond return resized_cond
def set_step(self, timestep: torch.Tensor, model_options: dict[str]):
indexes = torch.where(model_options["transformer_options"]["sample_sigmas"] == timestep[0])
self._step = int(indexes[0])
def get_context_windows(self, model: BaseModel, x_in: torch.Tensor, model_options: dict[str]) -> list[IndexListContextWindow]: def get_context_windows(self, model: BaseModel, x_in: torch.Tensor, model_options: dict[str]) -> list[IndexListContextWindow]:
full_length = x_in.size(self.dim) # TODO: choose dim based on model full_length = x_in.size(self.dim) # TODO: choose dim based on model
context_windows = get_context_windows(full_length, self, model_options) context_windows = self.context_schedule.func(full_length, self, model_options)
context_windows = [IndexListContextWindow(window, dim=self.dim) for window in context_windows] context_windows = [IndexListContextWindow(window, dim=self.dim) for window in context_windows]
return context_windows return context_windows
def execute(self, calc_cond_batch: Callable, model: BaseModel, conds: list[list[dict]], x_in: torch.Tensor, timestep: torch.Tensor, model_options: dict[str]): def execute(self, calc_cond_batch: Callable, model: BaseModel, conds: list[list[dict]], x_in: torch.Tensor, timestep: torch.Tensor, model_options: dict[str]):
self.set_step(timestep, model_options)
context_windows = self.get_context_windows(model, x_in, model_options) context_windows = self.get_context_windows(model, x_in, model_options)
enumerated_context_windows = list(enumerate(context_windows)) enumerated_context_windows = list(enumerate(context_windows))
conds_final = [torch.zeros_like(x_in) for _ in conds] conds_final = [torch.zeros_like(x_in) for _ in conds]
if self.fuse_method == ContextFuseMethod.RELATIVE: if self.fuse_method.name == ContextFuseMethods.RELATIVE:
counts_final = [torch.ones(get_shape_for_dim(x_in, self.dim), device=x_in.device) for _ in conds] counts_final = [torch.ones(get_shape_for_dim(x_in, self.dim), device=x_in.device) for _ in conds]
else: else:
counts_final = [torch.zeros(get_shape_for_dim(x_in, self.dim), device=x_in.device) for _ in conds] counts_final = [torch.zeros(get_shape_for_dim(x_in, self.dim), device=x_in.device) for _ in conds]
biases_final = [([0.0] * x_in.shape[self.dim]) for _ in conds] biases_final = [([0.0] * x_in.shape[self.dim]) for _ in conds]
# TODO: add callback here
for enum_window in enumerated_context_windows: for enum_window in enumerated_context_windows:
results = self.evaluate_context_windows(calc_cond_batch, model, x_in, conds, timestep, [enum_window], model_options) results = self.evaluate_context_windows(calc_cond_batch, model, x_in, conds, timestep, [enum_window], model_options)
for result in results: for result in results:
self.combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.window, result.window_idx, len(enumerated_context_windows), timestep, self.combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.window, result.window_idx, len(enumerated_context_windows), timestep,
conds_final, counts_final, biases_final) conds_final, counts_final, biases_final)
# finalize conds # finalize conds
if self.fuse_method == ContextFuseMethod.RELATIVE: if self.fuse_method.name == ContextFuseMethods.RELATIVE:
# relative is already normalized, so return as is # relative is already normalized, so return as is
del counts_final del counts_final
return conds_final return conds_final
@ -189,36 +205,6 @@ class IndexListContextHandler(ContextHandlerABC):
for window_idx, window in enumerated_context_windows: for window_idx, window in enumerated_context_windows:
# allow processing to end between context window executions for faster Cancel # allow processing to end between context window executions for faster Cancel
comfy.model_management.throw_exception_if_processing_interrupted() comfy.model_management.throw_exception_if_processing_interrupted()
# if device is None:
# # non-MultiGPU execution
# comfy.model_management.throw_exception_if_processing_interrupted()
# elif first_device == device:
# # MultiGPU execution
# try:
# comfy.model_management.throw_exception_if_processing_interrupted()
# if ADGS.is_processing_interrupted():
# break
# except comfy.model_management.InterruptProcessingException:
# ADGS.interrupt_processing()
# break
# else:
# # MultiGPU execution
# if ADGS.is_processing_interrupted():
# break
# ADGS.params.sub_idxs = ctx_idxs
# if device is None:
# motion_models_devices = ADGS.motion_models_devices.values()
# else:
# motion_models_devices = ADGS.motion_models_devices.get(device, None)
# if motion_models_devices is None:
# motion_models_devices = []
# else:
# motion_models_devices = [motion_models_devices]
# model = ADGS.model_patcher_devices[device].model
# for motion_models in motion_models_devices:
# motion_models.set_sub_idxs(ctx_idxs)
# motion_models.set_video_length(len(ctx_idxs), ADGS.params.full_length)
# TODO: add callback here # TODO: add callback here
@ -240,7 +226,7 @@ class IndexListContextHandler(ContextHandlerABC):
def combine_context_window_results(self, x_in: torch.Tensor, sub_conds_out, sub_conds, window: IndexListContextWindow, window_idx: int, total_windows: int, timestep: torch.Tensor, def combine_context_window_results(self, x_in: torch.Tensor, sub_conds_out, sub_conds, window: IndexListContextWindow, window_idx: int, total_windows: int, timestep: torch.Tensor,
conds_final: list[torch.Tensor], counts_final: list[torch.Tensor], biases_final: list[torch.Tensor]): conds_final: list[torch.Tensor], counts_final: list[torch.Tensor], biases_final: list[torch.Tensor]):
if self.fuse_method == ContextFuseMethod.RELATIVE: if self.fuse_method.name == ContextFuseMethods.RELATIVE:
for pos, idx in enumerate(window.index_list): for pos, idx in enumerate(window.index_list):
# bias is the influence of a specific index in relation to the whole context window # bias is the influence of a specific index in relation to the whole context window
bias = 1 - abs(idx - (window.index_list[0] + window.index_list[-1]) / 2) / ((window.index_list[-1] - window.index_list[0] + 1e-2) / 2) bias = 1 - abs(idx - (window.index_list[0] + window.index_list[-1]) / 2) / ((window.index_list[-1] - window.index_list[0] + 1e-2) / 2)
@ -260,10 +246,8 @@ class IndexListContextHandler(ContextHandlerABC):
for i in range(len(sub_conds_out)): for i in range(len(sub_conds_out)):
window.add_window(conds_final[i], sub_conds_out[i] * weights_tensor) window.add_window(conds_final[i], sub_conds_out[i] * weights_tensor)
window.add_window(counts_final[i], weights_tensor) window.add_window(counts_final[i], weights_tensor)
# # handle NaiveReuse
# NAIVE.cache_first_context_results(window_idx, ctx_idxs, sub_conds, conds_final, counts_final) # TODO: add callback here
# # handle ContextRef
# CREF.finalize_step()
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)
@ -301,9 +285,9 @@ def create_windows_uniform_looped(num_frames: int, handler: IndexListContextHand
context_stride = min(handler.context_stride, int(np.ceil(np.log2(num_frames / handler.context_length))) + 1) context_stride = min(handler.context_stride, int(np.ceil(np.log2(num_frames / handler.context_length))) + 1)
# obtain uniform windows as normal, looping and all # obtain uniform windows as normal, looping and all
for context_step in 1 << np.arange(context_stride): for context_step in 1 << np.arange(context_stride):
pad = int(round(num_frames * ordered_halving(handler.step))) pad = int(round(num_frames * ordered_halving(handler._step)))
for j in range( for j in range(
int(ordered_halving(handler.step) * context_step) + pad, int(ordered_halving(handler._step) * context_step) + pad,
num_frames + pad + (0 if handler.closed_loop else -handler.context_overlap), num_frames + pad + (0 if handler.closed_loop else -handler.context_overlap),
(handler.context_length * context_step - handler.context_overlap), (handler.context_length * context_step - handler.context_overlap),
): ):
@ -323,9 +307,9 @@ def create_windows_uniform_standard(num_frames: int, handler: IndexListContextHa
context_stride = min(handler.context_stride, int(np.ceil(np.log2(num_frames / handler.context_length))) + 1) context_stride = min(handler.context_stride, int(np.ceil(np.log2(num_frames / handler.context_length))) + 1)
# first, obtain uniform windows as normal, looping and all # first, obtain uniform windows as normal, looping and all
for context_step in 1 << np.arange(context_stride): for context_step in 1 << np.arange(context_stride):
pad = int(round(num_frames * ordered_halving(handler.step))) pad = int(round(num_frames * ordered_halving(handler._step)))
for j in range( for j in range(
int(ordered_halving(handler.step) * context_step) + pad, int(ordered_halving(handler._step) * context_step) + pad,
num_frames + pad + (-handler.context_overlap), num_frames + pad + (-handler.context_overlap),
(handler.context_length * context_step - handler.context_overlap), (handler.context_length * context_step - handler.context_overlap),
): ):
@ -395,13 +379,6 @@ def create_windows_default(num_frames: int, handler: IndexListContextHandler):
return [list(range(num_frames))] return [list(range(num_frames))]
def get_context_windows(num_frames: int, handler: IndexListContextHandler, model_options: dict[str]) -> list[list[int]]:
context_func = CONTEXT_MAPPING.get(handler.context_schedule, None)
if not context_func:
raise ValueError(f"Unknown context_schedule '{handler.context_schedule}'.")
return context_func(num_frames, handler, model_options)
CONTEXT_MAPPING = { CONTEXT_MAPPING = {
ContextSchedules.UNIFORM_LOOPED: create_windows_uniform_looped, ContextSchedules.UNIFORM_LOOPED: create_windows_uniform_looped,
ContextSchedules.UNIFORM_STANDARD: create_windows_uniform_standard, ContextSchedules.UNIFORM_STANDARD: create_windows_uniform_standard,
@ -409,11 +386,16 @@ CONTEXT_MAPPING = {
ContextSchedules.BATCHED: create_windows_batched, ContextSchedules.BATCHED: create_windows_batched,
} }
def get_matching_context_schedule(context_schedule: str) -> ContextSchedule:
func = CONTEXT_MAPPING.get(context_schedule, None)
if func is None:
raise ValueError(f"Unknown context_schedule '{context_schedule}'.")
return ContextSchedule(context_schedule, func)
def get_context_weights(length: int, full_length: int, idxs: list[int], handler: IndexListContextHandler, sigma: torch.Tensor=None): def get_context_weights(length: int, full_length: int, idxs: list[int], handler: IndexListContextHandler, sigma: torch.Tensor=None):
weights_func = FUSE_MAPPING.get(handler.fuse_method, None) return handler.fuse_method.func(length, sigma=sigma, handler=handler, full_length=full_length, idxs=idxs)
if not weights_func:
raise ValueError(f"Unknown fuse_method '{handler.fuse_method}'.")
return weights_func(length, sigma=sigma, handler=handler, full_length=full_length, idxs=idxs)
def create_weights_flat(length: int, **kwargs) -> list[float]: def create_weights_flat(length: int, **kwargs) -> list[float]:
@ -445,7 +427,7 @@ def create_weights_overlap_linear(length: int, full_length: int, idxs: list[int]
weights_torch[-handler.context_overlap:] = ramp_down weights_torch[-handler.context_overlap:] = ramp_down
return weights_torch return weights_torch
class ContextFuseMethod: class ContextFuseMethods:
FLAT = "flat" FLAT = "flat"
PYRAMID = "pyramid" PYRAMID = "pyramid"
RELATIVE = "relative" RELATIVE = "relative"
@ -456,12 +438,18 @@ class ContextFuseMethod:
FUSE_MAPPING = { FUSE_MAPPING = {
ContextFuseMethod.FLAT: create_weights_flat, ContextFuseMethods.FLAT: create_weights_flat,
ContextFuseMethod.PYRAMID: create_weights_pyramid, ContextFuseMethods.PYRAMID: create_weights_pyramid,
ContextFuseMethod.RELATIVE: create_weights_pyramid, ContextFuseMethods.RELATIVE: create_weights_pyramid,
ContextFuseMethod.OVERLAP_LINEAR: create_weights_overlap_linear, ContextFuseMethods.OVERLAP_LINEAR: create_weights_overlap_linear,
} }
def get_matching_fuse_method(fuse_method: str) -> ContextFuseMethod:
func = FUSE_MAPPING.get(fuse_method, None)
if func is None:
raise ValueError(f"Unknown fuse_method '{fuse_method}'.")
return ContextFuseMethod(fuse_method, func)
# Returns fraction that has denominator that is a power of 2 # Returns fraction that has denominator that is a power of 2
def ordered_halving(val): def ordered_halving(val):
# get binary value, padded with 0s for 64 bits # get binary value, padded with 0s for 64 bits

View File

@ -70,10 +70,14 @@ class ContextWindowsNode(io.ComfyNode):
io.Int.Input("context_overlap", min=0, default=0, tooltip="The overlap of the context window."), io.Int.Input("context_overlap", min=0, default=0, tooltip="The overlap of the context window."),
io.Combo.Input("context_schedule", options=[ io.Combo.Input("context_schedule", options=[
comfy.context_windows.ContextSchedules.STATIC_STANDARD, comfy.context_windows.ContextSchedules.STATIC_STANDARD,
comfy.context_windows.ContextSchedules.UNIFORM_STANDARD,
comfy.context_windows.ContextSchedules.UNIFORM_LOOPED,
comfy.context_windows.ContextSchedules.BATCHED, comfy.context_windows.ContextSchedules.BATCHED,
], tooltip="The stride of the context window."), ], tooltip="The stride of the context window."),
io.Combo.Input("fuse_method", options=comfy.context_windows.ContextFuseMethod.LIST_STATIC,default=comfy.context_windows.ContextFuseMethod.PYRAMID, tooltip="The method to use to fuse the context windows."), io.Int.Input("context_stride", min=1, default=1, tooltip="The stride of the context window."),
io.Int.Input("dim", min=0, max=2, default=0, tooltip="The dimension to apply the context windows to."), io.Boolean.Input("closed_loop", default=False, tooltip="Whether to close the context window loop."),
io.Combo.Input("fuse_method", options=comfy.context_windows.ContextFuseMethods.LIST_STATIC, default=comfy.context_windows.ContextFuseMethods.PYRAMID, tooltip="The method to use to fuse the context windows."),
io.Int.Input("dim", min=0, max=5, default=0, tooltip="The dimension to apply the context windows to."),
], ],
outputs=[ outputs=[
io.Model.Output(tooltip="The model with context windows applied during sampling."), io.Model.Output(tooltip="The model with context windows applied during sampling."),
@ -82,13 +86,15 @@ class ContextWindowsNode(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, model: io.Model.Type, context_length: int, context_overlap: int, context_schedule: str, fuse_method: str, dim: int) -> io.Model: def execute(cls, model: io.Model.Type, context_length: int, context_overlap: int, context_schedule: str, context_stride: int, closed_loop: bool, fuse_method: str, dim: int) -> io.Model:
model = model.clone() model = model.clone()
model.model_options["context_handler"] = comfy.context_windows.IndexListContextHandler( model.model_options["context_handler"] = comfy.context_windows.IndexListContextHandler(
context_schedule=context_schedule, context_schedule=comfy.context_windows.get_matching_context_schedule(context_schedule),
fuse_method=fuse_method, fuse_method=comfy.context_windows.get_matching_fuse_method(fuse_method),
context_length=context_length, context_length=context_length,
context_overlap=context_overlap, context_overlap=context_overlap,
context_stride=context_stride,
closed_loop=closed_loop,
dim=dim) dim=dim)
create_prepare_sampling_wrapper(model.model_options) create_prepare_sampling_wrapper(model.model_options)
#create_outer_sampler_wrapper(model.model_options) #create_outer_sampler_wrapper(model.model_options)
@ -116,10 +122,10 @@ class WanContextWindowsNode(ContextWindowsNode):
return schema return schema
@classmethod @classmethod
def execute(cls, model: io.Model.Type, context_length: int, context_overlap: int, context_schedule: str, fuse_method: str) -> io.Model: def execute(cls, model: io.Model.Type, context_length: int, context_overlap: int, context_schedule: str, context_stride: int, closed_loop: bool, fuse_method: str) -> io.Model:
context_length = max(((context_length - 1) // 4) + 1, 1) # at least length 1 context_length = max(((context_length - 1) // 4) + 1, 1) # at least length 1
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, 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]]: