relative fuse_method only exposed for Standard Static method, added commented out view options version of relative fuse_method - not working as intended, I'll need to see if it's possible to make it work there later
This commit is contained in:
@@ -12,7 +12,8 @@ class ContextFuseMethod:
|
||||
PYRAMID = "pyramid"
|
||||
RELATIVE = "relative"
|
||||
|
||||
LIST = [PYRAMID, RELATIVE, FLAT]
|
||||
LIST = [PYRAMID, FLAT]
|
||||
LIST_STATIC = [PYRAMID, RELATIVE, FLAT]
|
||||
|
||||
|
||||
class ContextType:
|
||||
@@ -332,6 +333,7 @@ def create_weights_pyramid(length: int, **kwargs) -> list[float]:
|
||||
FUSE_MAPPING = {
|
||||
ContextFuseMethod.FLAT: create_weights_flat,
|
||||
ContextFuseMethod.PYRAMID: create_weights_pyramid,
|
||||
ContextFuseMethod.RELATIVE: create_weights_pyramid,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ from comfy.utils import repeat_to_batch_size
|
||||
import comfy.ops
|
||||
import comfy.model_management
|
||||
|
||||
from .context import ContextOptions, get_context_weights, get_context_windows
|
||||
from .context import ContextFuseMethod, ContextOptions, get_context_weights, get_context_windows
|
||||
from .utils_motion import CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch
|
||||
from .utils_model import BetaSchedules, ModelTypeSD
|
||||
from .logger import logger
|
||||
@@ -813,11 +813,9 @@ class TemporalTransformerBlock(nn.Module):
|
||||
hidden_states = rearrange(hidden_states, "(b f) d c -> b f d c", f=video_length)
|
||||
value_final = torch.zeros_like(hidden_states)
|
||||
count_final = torch.zeros_like(hidden_states)
|
||||
# bias_final = [0.0] * video_length
|
||||
batched_conds = hidden_states.size(1) // video_length
|
||||
for sub_idxs in views:
|
||||
weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds
|
||||
weights_tensor = torch.Tensor(weights).to(device=hidden_states.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
|
||||
sub_hidden_states = rearrange(hidden_states[:, sub_idxs], "b f d c -> (b f) d c")
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(sub_hidden_states).to(sub_hidden_states.dtype)
|
||||
@@ -834,14 +832,35 @@ class TemporalTransformerBlock(nn.Module):
|
||||
)
|
||||
sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs))
|
||||
|
||||
# if view_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
# for pos, idx in enumerate(sub_idxs):
|
||||
# # bias is the influence of a specific index in relation to the whole context window
|
||||
# bias = 1 - abs(idx - (sub_idxs[0] + sub_idxs[-1]) / 2) / ((sub_idxs[-1] - sub_idxs[0] + 1e-2) / 2)
|
||||
# bias = max(1e-2, bias)
|
||||
# # take weighted averate relative to total bias of current idx
|
||||
# bias_total = bias_final[idx]
|
||||
# prev_weight = torch.tensor([bias_total / (bias_total + bias)],
|
||||
# dtype=value_final.dtype, device=value_final.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
# #prev_weight = torch.cat([prev_weight]*value_final.shape[1], dim=1)
|
||||
# new_weight = torch.tensor([bias / (bias_total + bias)],
|
||||
# dtype=value_final.dtype, device=value_final.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
# #new_weight = torch.cat([new_weight]*value_final.shape[1], dim=1)
|
||||
# test = value_final[:, idx:idx+1, :, :]
|
||||
# value_final[:, idx:idx+1, :, :] = value_final[:, idx:idx+1, :, :] * prev_weight + sub_hidden_states[:, pos:pos+1, : ,:] * new_weight
|
||||
# bias_final[idx] = bias_total + bias
|
||||
# else:
|
||||
weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds
|
||||
weights_tensor = torch.Tensor(weights).to(device=hidden_states.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
value_final[:, sub_idxs] += sub_hidden_states * weights_tensor
|
||||
count_final[:, sub_idxs] += weights_tensor
|
||||
|
||||
# get weighted average of sub_hidden_states
|
||||
# get weighted average of sub_hidden_states, if fuse method requires it
|
||||
# if view_options.fuse_method != ContextFuseMethod.RELATIVE:
|
||||
hidden_states = value_final / count_final
|
||||
hidden_states = rearrange(hidden_states, "b f d c -> (b f) d c")
|
||||
del value_final
|
||||
del count_final
|
||||
# del bias_final
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
|
||||
@@ -145,7 +145,7 @@ class StandardStaticContextOptionsNode:
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST,),
|
||||
"fuse_method": (ContextFuseMethod.LIST_STATIC,),
|
||||
"use_on_equal_length": ("BOOLEAN", {"default": False},),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
|
||||
@@ -12,7 +12,7 @@ from PIL.PngImagePlugin import PngInfo
|
||||
import folder_paths
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import ContextSchedules, ContextOptions
|
||||
from .context import ContextOptionsGroup, ContextOptions, ContextSchedules
|
||||
from .logger import logger
|
||||
from .utils_model import Folders, BetaSchedules, get_available_motion_models
|
||||
from .model_injection import ModelPatcherAndInjector, InjectionParams, MotionModelGroup, load_motion_module_gen1
|
||||
@@ -109,16 +109,18 @@ class AnimateDiffLoaderAdvanced_Deprecated:
|
||||
model_name=model_name,
|
||||
apply_v2_properly=False,
|
||||
)
|
||||
# set context settings
|
||||
params.set_context(
|
||||
context_group = ContextOptionsGroup()
|
||||
context_group.add(
|
||||
ContextOptions(
|
||||
context_length=context_length,
|
||||
context_stride=context_stride,
|
||||
context_overlap=context_overlap,
|
||||
context_schedule=context_schedule,
|
||||
closed_loop=closed_loop,
|
||||
)
|
||||
)
|
||||
)
|
||||
# set context settings
|
||||
params.set_context(context_options=context_group)
|
||||
# inject for use in sampling code
|
||||
model = ModelPatcherAndInjector(model)
|
||||
model.motion_models = MotionModelGroup(motion_model)
|
||||
|
||||
@@ -492,7 +492,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
bias_total = bias_final[(full_length*n)+idx]
|
||||
prev_weight = (bias_total / (bias_total + bias))
|
||||
new_weight = (bias / (bias_total + bias))
|
||||
cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight
|
||||
cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight
|
||||
uncond_final[(full_length*n)+idx] = uncond_final[(full_length*n)+idx] * prev_weight + sub_uncond_out[(full_length*n)+pos] * new_weight
|
||||
bias_final[(full_length*n)+idx] = bias_total + bias
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user