Added overlap-linear fuse method, marked some fuse methods as experimental to make the dropdown less overwhelming, made sigma-related fuse methods work with View Options

This commit is contained in:
Jedrzej Kosinski
2025-03-08 23:36:52 -06:00
parent 0542cd97c3
commit 599b71807d
3 changed files with 30 additions and 15 deletions
+24 -6
View File
@@ -20,6 +20,8 @@ class ContextFuseMethod:
FLAT = "flat"
PYRAMID = "pyramid"
RELATIVE = "relative"
OVERLAP_LINEAR = "overlap-linear"
RANDOM = "🔬random"
RANDOM_DEPR = "random"
GAUSS_SIGMA = "🔬gauss-sigma"
@@ -33,8 +35,8 @@ class ContextFuseMethod:
PYRAMID_SIGMA_INV = "🔬pyramid-sigma inverse"
PYRAMID_SIGMA_INV_DEPR = "pyramid-sigma inverse"
LIST = [PYRAMID, FLAT, DELAYED_REVERSE_SAWTOOTH, PYRAMID_SIGMA, PYRAMID_SIGMA_INV, GAUSS_SIGMA, GAUSS_SIGMA_INV, RANDOM]
LIST_STATIC = [PYRAMID, RELATIVE, FLAT, DELAYED_REVERSE_SAWTOOTH, PYRAMID_SIGMA, PYRAMID_SIGMA_INV, GAUSS_SIGMA, GAUSS_SIGMA_INV, RANDOM]
LIST = [PYRAMID, FLAT, OVERLAP_LINEAR, DELAYED_REVERSE_SAWTOOTH, PYRAMID_SIGMA, PYRAMID_SIGMA_INV, GAUSS_SIGMA, GAUSS_SIGMA_INV, RANDOM]
LIST_STATIC = [PYRAMID, RELATIVE, FLAT, OVERLAP_LINEAR, DELAYED_REVERSE_SAWTOOTH, PYRAMID_SIGMA, PYRAMID_SIGMA_INV, GAUSS_SIGMA, GAUSS_SIGMA_INV, RANDOM]
class ContextType:
@@ -360,11 +362,11 @@ CONTEXT_MAPPING = {
}
def get_context_weights(num_frames: int, fuse_method: str, sigma: Tensor = None):
weights_func = FUSE_MAPPING.get(fuse_method, None)
def get_context_weights(length: int, full_length: int, idxs: list[int], ctx_opts: ContextOptions, sigma: Tensor=None):
weights_func = FUSE_MAPPING.get(ctx_opts.fuse_method, None)
if not weights_func:
raise ValueError(f"Unknown fuse_method '{fuse_method}'.")
return weights_func(num_frames, sigma=sigma )
raise ValueError(f"Unknown fuse_method '{ctx_opts.fuse_method}'.")
return weights_func(length, sigma=sigma, ctx_opts=ctx_opts, full_length=full_length, idxs=idxs)
def create_weights_flat(length: int, **kwargs) -> list[float]:
@@ -382,6 +384,20 @@ def create_weights_pyramid(length: int, **kwargs) -> list[float]:
weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1))
return weight_sequence
def create_weights_overlap_linear(length: int, full_length: int, idxs: list[int], ctx_opts: ContextOptions, **kwargs):
# based on code in Kijai's WanVideoWrapper: https://github.com/kijai/ComfyUI-WanVideoWrapper/blob/dbb2523b37e4ccdf45127e5ae33e31362f755c8e/nodes.py#L1302
# only expected overlap is given different weights
weights_torch = torch.ones((length))
# blend left-side on all except first window
if min(idxs) > 0:
ramp_up = torch.linspace(1e-37, 1, ctx_opts.context_overlap)
weights_torch[:ctx_opts.context_overlap] = ramp_up
# blend right-side on all except last window
if max(idxs) < full_length-1:
ramp_down = torch.linspace(1, 1e-37, ctx_opts.context_overlap)
weights_torch[-ctx_opts.context_overlap:] = ramp_down
return weights_torch
def create_weights_random(length: int, **kwargs) -> list[float]:
if length % 2 == 0:
max_weight = length // 2
@@ -460,6 +476,8 @@ FUSE_MAPPING = {
ContextFuseMethod.FLAT: create_weights_flat,
ContextFuseMethod.PYRAMID: create_weights_pyramid,
ContextFuseMethod.RELATIVE: create_weights_pyramid,
ContextFuseMethod.OVERLAP_LINEAR: create_weights_overlap_linear,
# experimental
ContextFuseMethod.GAUSS_SIGMA: create_weights_gauss_sigma,
ContextFuseMethod.GAUSS_SIGMA_DEPR: create_weights_gauss_sigma,
ContextFuseMethod.GAUSS_SIGMA_INV: create_weights_gauss_sigma_inv,
+1 -1
View File
@@ -1387,7 +1387,7 @@ class TemporalTransformerBlock(nn.Module):
count += 1
sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs))
weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds
weights = get_context_weights(len(sub_idxs), video_length, sub_idxs, view_options, sigma=transformer_options["sigmas"]) * 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
+5 -8
View File
@@ -783,9 +783,6 @@ def sliding_calc_cond_batch(executor: Callable, model, conds: list[list[dict]],
multigpu_windows = {}
start_idx = 0
for device, work in ctxs_relative_work.items():
# if device == x_in.device:
# continue
# multigpu_windows[device] = enumerated_context_windows
if work == 0:
continue
end_idx = start_idx + work
@@ -817,14 +814,14 @@ def sliding_calc_cond_batch(executor: Callable, model, conds: list[list[dict]],
for results in combined_results:
for result in results:
combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.ctx_idxs, result.window_idx, timestep,
combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.ctx_idxs, result.window_idx, len(enumerated_context_windows), timestep,
ADGS, NAIVE, CREF, conds_final, counts_final, biases_final)
else:
for enum_window in enumerated_context_windows:
results = evaluate_context_windows(executor, model, x_in, conds, timestep, [enum_window], model_options, CREF, ADGS)
for result in results:
combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.ctx_idxs, result.window_idx, timestep,
combine_context_window_results(x_in, result.sub_conds_out, result.sub_conds, result.ctx_idxs, result.window_idx, len(enumerated_context_windows), timestep,
ADGS, NAIVE, CREF, conds_final, counts_final, biases_final)
finally:
CREF.cleanup(model_options)
@@ -834,7 +831,7 @@ def sliding_calc_cond_batch(executor: Callable, model, conds: list[list[dict]],
# finalize conds
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
# already normalized, so return as is
# relative is already normalized, so return as is
del counts_final
return conds_final
else:
@@ -898,7 +895,7 @@ def evaluate_context_windows(executor, model: BaseModel, x_in: Tensor, conds, ti
return results
def combine_context_window_results(x_in: Tensor, sub_conds_out, sub_conds, ctx_idxs: list[int], window_idx: int, timestep,
def combine_context_window_results(x_in: Tensor, sub_conds_out, sub_conds, ctx_idxs: list[int], window_idx: int, total_windows: int, timestep,
ADGS: AnimateDiffGlobalState, NAIVE: NaiveReuseHandler, CREF: ContextRefHandler,
conds_final: list[Tensor], counts_final: list[Tensor], biases_final: list[Tensor]):
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
@@ -915,7 +912,7 @@ def combine_context_window_results(x_in: Tensor, sub_conds_out, sub_conds, ctx_i
biases_final[i][idx] = bias_total + bias
else:
# add conds and counts based on weights of fuse method
weights = get_context_weights(len(ctx_idxs), ADGS.params.context_options.fuse_method, sigma=timestep)
weights = get_context_weights(len(ctx_idxs), x_in.shape[0], ctx_idxs, ADGS.params.context_options, sigma=timestep)
weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
for i in range(len(sub_conds_out)):
conds_final[i][ctx_idxs] += sub_conds_out[i] * weights_tensor