mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-15 22:56:49 +08:00
Made context windows compatible with different dimensions; works for WAN, but results are bad
This commit is contained in:
parent
c0ff26bd16
commit
b114b5dc6b
@ -57,17 +57,27 @@ class IndexListContextWindow(ContextWindowABC):
|
||||
if dim is None:
|
||||
dim = self.dim
|
||||
if dim == 0:
|
||||
if full.shape[dim] == 1:
|
||||
return full
|
||||
return full[self.index_list].to(device)
|
||||
else:
|
||||
elif dim == 1:
|
||||
return full[:, self.index_list].to(device)
|
||||
elif dim == 2:
|
||||
return full[:, :, self.index_list].to(device)
|
||||
else:
|
||||
raise ValueError(f"Invalid dimension: {dim}")
|
||||
|
||||
def add_window(self, full: torch.Tensor, to_window: torch.Tensor, dim=None) -> torch.Tensor:
|
||||
if dim is None:
|
||||
dim = self.dim
|
||||
if dim == 0:
|
||||
full[self.index_list] += to_window
|
||||
else:
|
||||
elif dim == 1:
|
||||
full[:, self.index_list] += to_window
|
||||
elif dim == 2:
|
||||
full[:, :, self.index_list] += to_window
|
||||
else:
|
||||
raise ValueError(f"Invalid dimension: {dim}")
|
||||
return full
|
||||
|
||||
ContextResults = collections.namedtuple("ContextResults", ['window_idx', 'sub_conds_out', 'sub_conds', 'window'])
|
||||
@ -84,8 +94,8 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
|
||||
def should_use_context(self, model: BaseModel, conds: list[list[dict]], x_in: torch.Tensor, timestep: torch.Tensor, model_options: dict[str]) -> bool:
|
||||
# for now, assume first dim is batch - should have stored on BaseModel in actual implementation
|
||||
if x_in.size(0) > self.context_length:
|
||||
logging.info(f"Using context windows {self.context_length} for {x_in.size(0)} frames.")
|
||||
if x_in.size(self.dim) > self.context_length:
|
||||
logging.info(f"Using context windows {self.context_length} for {x_in.size(self.dim)} frames.")
|
||||
return True
|
||||
return False
|
||||
|
||||
@ -110,7 +120,7 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
cond_item = actual_cond[key]
|
||||
if isinstance(cond_item, torch.Tensor):
|
||||
# check that tensor is the expected length - x.size(0)
|
||||
if cond_item.size(0) == x_in.size(0):
|
||||
if cond_item.size(self.dim) == x_in.size(self.dim):
|
||||
# if so, it's subsetting time - tell controls the expected indeces so they can handle them
|
||||
actual_cond_item = window.get_tensor(cond_item)
|
||||
resized_actual_cond[key] = actual_cond_item.to(device)
|
||||
@ -124,11 +134,11 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
# when in dictionary, look for tensors and CONDCrossAttn [comfy/conds.py] (has cond attr that is a tensor)
|
||||
for cond_key, cond_value in new_cond_item.items():
|
||||
if isinstance(cond_value, torch.Tensor):
|
||||
if cond_value.size(0) == x_in.size(0):
|
||||
if cond_value.size(self.dim) == x_in.size(self.dim):
|
||||
new_cond_item[cond_key] = window.get_tensor(cond_value, device)
|
||||
# if has cond that is a Tensor, check if needs to be subset
|
||||
elif hasattr(cond_value, "cond") and isinstance(cond_value.cond, torch.Tensor):
|
||||
if cond_value.cond.size(0) == x_in.size(0):
|
||||
if cond_value.cond.size(self.dim) == x_in.size(self.dim):
|
||||
new_cond_item[cond_key] = cond_value._copy_with(window.get_tensor(cond_value.cond, device))
|
||||
elif cond_key == "num_video_frames": # for SVD
|
||||
new_cond_item[cond_key] = cond_value._copy_with(cond_value.cond)
|
||||
@ -142,7 +152,7 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
return resized_cond
|
||||
|
||||
def get_context_windows(self, model: BaseModel, x_in: torch.Tensor, model_options: dict[str]) -> list[IndexListContextWindow]:
|
||||
full_length = x_in.size(0) # 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 = [IndexListContextWindow(window, dim=self.dim) for window in context_windows]
|
||||
return context_windows
|
||||
@ -153,16 +163,16 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
|
||||
conds_final = [torch.zeros_like(x_in) for _ in conds]
|
||||
if self.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
counts_final = [torch.ones((x_in.shape[0], 1, 1, 1), 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:
|
||||
counts_final = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds]
|
||||
biases_final = [([0.0] * x_in.shape[0]) 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]
|
||||
|
||||
for enum_window in enumerated_context_windows:
|
||||
results = self.evaluate_context_windows(calc_cond_batch, model, x_in, conds, timestep, [enum_window], model_options)
|
||||
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,
|
||||
conds_final, counts_final, biases_final)
|
||||
results = self.evaluate_context_windows(calc_cond_batch, model, x_in, conds, timestep, [enum_window], model_options)
|
||||
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,
|
||||
conds_final, counts_final, biases_final)
|
||||
# finalize conds
|
||||
if self.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
# relative is already normalized, so return as is
|
||||
@ -246,8 +256,8 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
biases_final[i][idx] = bias_total + bias
|
||||
else:
|
||||
# add conds and counts based on weights of fuse method; TODO: account for dim not being 0
|
||||
weights = get_context_weights(window.context_length, x_in.shape[0], window.index_list, self, sigma=timestep)
|
||||
weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
|
||||
weights = get_context_weights(window.context_length, x_in.shape[self.dim], window.index_list, self, sigma=timestep)
|
||||
weights_tensor = match_weights_to_dim(weights, x_in, self.dim, device=x_in.device)
|
||||
for i in range(len(sub_conds_out)):
|
||||
window.add_window(conds_final[i], sub_conds_out[i] * weights_tensor)
|
||||
window.add_window(counts_final[i], weights_tensor)
|
||||
@ -256,7 +266,24 @@ class IndexListContextHandler(ContextHandlerABC):
|
||||
# # handle ContextRef
|
||||
# CREF.finalize_step()
|
||||
|
||||
def match_weights_to_dim(weights: list[float], x_in: torch.Tensor, dim: int, device=None) -> torch.Tensor:
|
||||
total_dims = len(x_in.shape)
|
||||
weights_tensor = torch.Tensor(weights).to(device=device)
|
||||
for _ in range(dim):
|
||||
weights_tensor = weights_tensor.unsqueeze(0)
|
||||
for _ in range(total_dims - dim - 1):
|
||||
weights_tensor = weights_tensor.unsqueeze(-1)
|
||||
return weights_tensor
|
||||
|
||||
def get_shape_for_dim(x_in: torch.Tensor, dim: int) -> list[int]:
|
||||
total_dims = len(x_in.shape)
|
||||
shape = []
|
||||
for _ in range(dim):
|
||||
shape.append(1)
|
||||
shape.append(x_in.shape[dim])
|
||||
for _ in range(total_dims - dim - 1):
|
||||
shape.append(1)
|
||||
return shape
|
||||
|
||||
class ContextSchedules:
|
||||
UNIFORM_LOOPED = "looped_uniform"
|
||||
|
||||
@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
import comfy.context_windows
|
||||
import comfy.patcher_extension
|
||||
import comfy.samplers
|
||||
import torch
|
||||
|
||||
|
||||
@ -13,7 +14,8 @@ def _prepare_sampling_wrapper(executor, model, noise_shape: torch.Tensor, *args,
|
||||
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:])
|
||||
noise_shape = list(noise_shape)
|
||||
noise_shape[handler.dim] = min(noise_shape[handler.dim], handler.context_length)
|
||||
return executor(model, noise_shape, *args, **kwargs)
|
||||
|
||||
|
||||
@ -23,6 +25,35 @@ def create_prepare_sampling_wrapper(model_options: dict):
|
||||
_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
|
||||
@ -41,7 +72,7 @@ class ContextWindowsNode(io.ComfyNode):
|
||||
comfy.context_windows.ContextSchedules.BATCHED,
|
||||
], 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("dim", min=0, max=1, default=0, tooltip="The dimension to apply the context windows to."),
|
||||
io.Int.Input("dim", min=0, max=2, default=0, tooltip="The dimension to apply the context windows to."),
|
||||
],
|
||||
outputs=[
|
||||
io.Model.Output(tooltip="The model with context windows applied during sampling."),
|
||||
@ -59,9 +90,21 @@ class ContextWindowsNode(io.ComfyNode):
|
||||
context_overlap=context_overlap,
|
||||
dim=dim)
|
||||
create_prepare_sampling_wrapper(model.model_options)
|
||||
#create_outer_sampler_wrapper(model.model_options)
|
||||
return io.NodeOutput(model)
|
||||
|
||||
|
||||
class WanContextWindowsNode(ContextWindowsNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
schema = super().define_schema()
|
||||
schema.node_id = "WanContextWindowsTest"
|
||||
schema.display_name = "Wan Context Windows Test"
|
||||
schema.description = "Test node for context windows (WAN)"
|
||||
schema.inputs.append(io.Int.Input("dim", min=0, max=2, default=0, tooltip="The dimension to apply the context windows to."))
|
||||
return schema
|
||||
|
||||
|
||||
class ContextWindowsExtension(ComfyExtension):
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user