Fix something stupid and start implementing context schedule

This commit is contained in:
kijai
2025-02-28 15:56:19 +02:00
parent dfe8000e63
commit 7a6418a953
4 changed files with 300 additions and 31 deletions
+184
View File
@@ -0,0 +1,184 @@
import numpy as np
from typing import Callable, Optional, List
def ordered_halving(val):
bin_str = f"{val:064b}"
bin_flip = bin_str[::-1]
as_int = int(bin_flip, 2)
return as_int / (1 << 64)
def does_window_roll_over(window: list[int], num_frames: int) -> tuple[bool, int]:
prev_val = -1
for i, val in enumerate(window):
val = val % num_frames
if val < prev_val:
return True, i
prev_val = val
return False, -1
def shift_window_to_start(window: list[int], num_frames: int):
start_val = window[0]
for i in range(len(window)):
# 1) subtract each element by start_val to move vals relative to the start of all frames
# 2) add num_frames and take modulus to get adjusted vals
window[i] = ((window[i] - start_val) + num_frames) % num_frames
def shift_window_to_end(window: list[int], num_frames: int):
# 1) shift window to start
shift_window_to_start(window, num_frames)
end_val = window[-1]
end_delta = num_frames - end_val - 1
for i in range(len(window)):
# 2) add end_delta to each val to slide windows to end
window[i] = window[i] + end_delta
def get_missing_indexes(windows: list[list[int]], num_frames: int) -> list[int]:
all_indexes = list(range(num_frames))
for w in windows:
for val in w:
try:
all_indexes.remove(val)
except ValueError:
pass
return all_indexes
def uniform_looped(
step: int = ...,
num_steps: Optional[int] = None,
num_frames: int = ...,
context_size: Optional[int] = None,
context_stride: int = 3,
context_overlap: int = 4,
closed_loop: bool = True,
):
if num_frames <= context_size:
yield list(range(num_frames))
return
context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1)
for context_step in 1 << np.arange(context_stride):
pad = int(round(num_frames * ordered_halving(step)))
for j in range(
int(ordered_halving(step) * context_step) + pad,
num_frames + pad + (0 if closed_loop else -context_overlap),
(context_size * context_step - context_overlap),
):
yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)]
#from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
def uniform_standard(
step: int = ...,
num_steps: Optional[int] = None,
num_frames: int = ...,
context_size: Optional[int] = None,
context_stride: int = 3,
context_overlap: int = 4,
closed_loop: bool = True,
):
windows = []
if num_frames <= context_size:
windows.append(list(range(num_frames)))
return windows
context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1)
for context_step in 1 << np.arange(context_stride):
pad = int(round(num_frames * ordered_halving(step)))
for j in range(
int(ordered_halving(step) * context_step) + pad,
num_frames + pad + (0 if closed_loop else -context_overlap),
(context_size * context_step - context_overlap),
):
windows.append([e % num_frames for e in range(j, j + context_size * context_step, context_step)])
# now that windows are created, shift any windows that loop, and delete duplicate windows
delete_idxs = []
win_i = 0
while win_i < len(windows):
# if window is rolls over itself, need to shift it
is_roll, roll_idx = does_window_roll_over(windows[win_i], num_frames)
if is_roll:
roll_val = windows[win_i][roll_idx] # roll_val might not be 0 for windows of higher strides
shift_window_to_end(windows[win_i], num_frames=num_frames)
# check if next window (cyclical) is missing roll_val
if roll_val not in windows[(win_i+1) % len(windows)]:
# need to insert new window here - just insert window starting at roll_val
windows.insert(win_i+1, list(range(roll_val, roll_val + context_size)))
# delete window if it's not unique
for pre_i in range(0, win_i):
if windows[win_i] == windows[pre_i]:
delete_idxs.append(win_i)
break
win_i += 1
# reverse delete_idxs so that they will be deleted in an order that doesn't break idx correlation
delete_idxs.reverse()
for i in delete_idxs:
windows.pop(i)
return windows
def static_standard(
step: int = ...,
num_steps: Optional[int] = None,
num_frames: int = ...,
context_size: Optional[int] = None,
context_stride: int = 3,
context_overlap: int = 4,
closed_loop: bool = True,
):
windows = []
if num_frames <= context_size:
windows.append(list(range(num_frames)))
return windows
# always return the same set of windows
delta = context_size - context_overlap
for start_idx in range(0, num_frames, delta):
# if past the end of frames, move start_idx back to allow same context_length
ending = start_idx + context_size
if ending >= num_frames:
final_delta = ending - num_frames
final_start_idx = start_idx - final_delta
windows.append(list(range(final_start_idx, final_start_idx + context_size)))
break
windows.append(list(range(start_idx, start_idx + context_size)))
return windows
def get_context_scheduler(name: str) -> Callable:
if name == "uniform_looped":
return uniform_looped
elif name == "uniform_standard":
return uniform_standard
elif name == "static_standard":
return static_standard
else:
raise ValueError(f"Unknown context_overlap policy {name}")
def get_total_steps(
scheduler,
timesteps: List[int],
num_steps: Optional[int] = None,
num_frames: int = ...,
context_size: Optional[int] = None,
context_stride: int = 3,
context_overlap: int = 4,
closed_loop: bool = True,
):
return sum(
len(
list(
scheduler(
i,
num_steps,
num_frames,
context_size,
context_stride,
context_overlap,
)
)
)
for i in range(len(timesteps))
)
+111 -15
View File
@@ -837,6 +837,36 @@ class WanVideoEmptyEmbeds:
#region Sampler
class WanVideoContextOptions:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"context_schedule": (["uniform_standard", "uniform_looped", "static_standard"],),
"context_frames": ("INT", {"default": 65, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of pixel frames in the context, NOTE: the latent space has 4 frames in 1"} ),
"context_stride": ("INT", {"default": 4, "min": 4, "max": 100, "step": 1, "tooltip": "Context stride as pixel frames, NOTE: the latent space has 4 frames in 1"} ),
"context_overlap": ("INT", {"default": 4, "min": 4, "max": 100, "step": 1, "tooltip": "Context overlap as pixel frames, NOTE: the latent space has 4 frames in 1"} ),
"freenoise": ("BOOLEAN", {"default": True, "tooltip": "Shuffle the noise"}),
}
}
RETURN_TYPES = ("WANVIDCONTEXT", )
RETURN_NAMES = ("context_options",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Context options for WanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow."
def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise):
context_options = {
"context_schedule":context_schedule,
"context_frames":context_frames,
"context_stride":context_stride,
"context_overlap":context_overlap,
"freenoise":freenoise
}
return (context_options,)
class WanVideoSampler:
@classmethod
def INPUT_TYPES(s):
@@ -862,6 +892,7 @@ class WanVideoSampler:
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"feta_args": ("FETAARGS", ),
"context_options": ("WANVIDCONTEXT", ),
}
}
@@ -871,7 +902,7 @@ class WanVideoSampler:
CATEGORY = "WanVideoWrapper"
def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index,
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0):
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None):
patcher = model
model = model.model
transformer = model.diffusion_model
@@ -940,16 +971,49 @@ class WanVideoSampler:
dtype=torch.float32,
device=torch.device("cpu"),
generator=seed_g)
latent_video_length = noise.shape[1]
if samples is not None:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * samples["samples"].squeeze(0).to(noise)
if context_options is not None:
context_schedule = context_options["context_schedule"]
context_frames = (context_options["context_frames"] - 1) // 4 + 1
context_stride = context_options["context_stride"] // 4
context_overlap = context_options["context_overlap"] // 4
if context_options["freenoise"]:
log.info("Applying FreeNoise")
# code and comments from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
delta = context_frames - context_overlap
for start_idx in range(0, latent_video_length-context_frames, delta):
place_idx = start_idx + context_frames
if place_idx >= latent_video_length:
break
end_idx = place_idx - 1
if end_idx + delta >= latent_video_length:
final_delta = latent_video_length - place_idx
list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
list_idx = list_idx[torch.randperm(final_delta, generator=seed_g)]
noise[:, place_idx:place_idx + final_delta, :, :] = noise[:, list_idx, :, :]
break
list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
list_idx = list_idx[torch.randperm(delta, generator=seed_g)]
noise[:, place_idx:place_idx + delta, :, :] = noise[:, list_idx, :, :]
log.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
from .context import get_context_scheduler
context = get_context_scheduler(context_schedule)
latent = noise.to(device)
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent.shape[2], k=riflex_freq_index),
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
@@ -1028,6 +1092,8 @@ class WanVideoSampler:
except:
pass
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*8}x{latent.shape[2]*8} with {steps} steps")
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True):
for i, t in enumerate(tqdm(timesteps)):
latent_model_input = [latent.to(device)]
@@ -1042,18 +1108,46 @@ class WanVideoSampler:
else:
disable_enhance()
#model inference start
noise_pred_cond = transformer(
latent_model_input, t=timestep, **arg_c)[0].to(offload_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
latent_model_input, t=timestep, **arg_null)[0].to(offload_device)
noise_pred = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
if context_options is not None:
counter = torch.zeros_like(latent_model_input[0], device=offload_device)
noise_pred = torch.zeros_like(latent_model_input[0], device=offload_device)
context_queue = list(context(
i, steps, latent_video_length, context_frames, context_stride, context_overlap,
))
for c in context_queue:
print(c)
partial_latent_model_input = [latent_model_input[0][:, c, :, :]]
print("partial_latent_model_input", partial_latent_model_input[0].shape)
#model inference start
noise_pred_cond = transformer(
partial_latent_model_input, t=timestep, **arg_c)[0].to(offload_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
partial_latent_model_input, t=timestep, **arg_null)[0].to(offload_device)
noise_pred_context = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
else:
noise_pred_context = noise_pred_cond
print(noise_pred.shape)
noise_pred[:, c, :, :] += noise_pred_context
noise_pred = noise_pred.float()
counter[:, c, :, :] += 1
#model inference end
noise_pred /= counter
else:
noise_pred = noise_pred_cond
#model inference end
#model inference start
noise_pred_cond = transformer(
latent_model_input, t=timestep, **arg_c)[0].to(offload_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
latent_model_input, t=timestep, **arg_null)[0].to(offload_device)
noise_pred = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
else:
noise_pred = noise_pred_cond
#model inference end
latent = latent.to(offload_device)
@@ -1281,7 +1375,8 @@ NODE_CLASS_MAPPINGS = {
"WanVideoEmptyEmbeds": WanVideoEmptyEmbeds,
"WanVideoLoraSelect": WanVideoLoraSelect,
"WanVideoLoraBlockEdit": WanVideoLoraBlockEdit,
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
"WanVideoContextOptions": WanVideoContextOptions
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -1301,5 +1396,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEmptyEmbeds": "WanVideo Empty Embeds",
"WanVideoLoraSelect": "WanVideo Lora Select",
"WanVideoLoraBlockEdit": "WanVideo Lora Block Edit",
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video"
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
"WanVideoContextOptions": "WanVideo Context Options"
}
+2 -12
View File
@@ -2,7 +2,6 @@
import math
import torch
import torch.cuda.amp as amp
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
@@ -29,7 +28,6 @@ def sinusoidal_embedding_1d(dim, position):
return x
@amp.autocast(enabled=False)
def rope_params(max_seq_len, dim, theta=10000, L_test=81, k=0):
assert dim % 2 == 0
freqs = torch.outer(
@@ -42,7 +40,8 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=81, k=0):
return freqs
@amp.autocast(enabled=False)
from comfy.model_management import get_torch_device, get_autocast_device
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply(x, grid_sizes, freqs):
n, c = x.size(2), x.size(3) // 2
@@ -332,8 +331,6 @@ class WanAttentionBlock(nn.Module):
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
assert e.dtype == torch.float32
#with amp.autocast(dtype=torch.float32):
# e = (self.modulation + e).chunk(6, dim=1)
e = (self.modulation.to(torch.float32) + e.to(torch.float32)).chunk(6, dim=1)
assert e[0].dtype == torch.float32
@@ -341,16 +338,12 @@ class WanAttentionBlock(nn.Module):
y = self.self_attn(
self.norm1(x).float() * (1 + e[1]) + e[0], seq_lens, grid_sizes,
freqs)
#with amp.autocast(dtype=torch.float32):
# x = x + y * e[2]
x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32))
# cross-attention & ffn function
def cross_attn_ffn(x, context, context_lens, e):
x = x + self.cross_attn(self.norm3(x), context, context_lens)
y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3])
#with amp.autocast(dtype=torch.float32):
# x = x + y * e[5]
x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32))
return x
@@ -382,9 +375,6 @@ class Head(nn.Module):
e(Tensor): Shape [B, C]
"""
assert e.dtype == torch.float32
# with amp.autocast(dtype=torch.float32):
# e = (self.modulation + e.unsqueeze(1)).chunk(2, dim=1)
# x = (self.head(self.norm(x) * (1 + e[1]) + e[0]))
e_unsqueezed = e.unsqueeze(1).to(torch.float32)
e = (self.modulation.to(torch.float32) + e_unsqueezed).chunk(2, dim=1)
normed = self.norm(x).to(torch.float32)
+3 -4
View File
@@ -2,11 +2,10 @@
import logging
import torch
import torch.cuda.amp as amp
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from comfy.model_management import get_torch_device, get_autocast_device
__all__ = [
'WanVAE',
]
@@ -648,14 +647,14 @@ class WanVAE:
"""
videos: A list of videos each with shape [C, T, H, W].
"""
with amp.autocast(dtype=self.dtype):
with torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False):
return [
self.model.encode(u.unsqueeze(0), self.scale).float().squeeze(0)
for u in videos
]
def decode(self, zs):
with amp.autocast(dtype=self.dtype):
with torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False):
return [
self.model.decode(u.unsqueeze(0),
self.scale).float().clamp_(-1, 1).squeeze(0)