From 7a6418a953c715fe875eb22c3db685bf358cd5b6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 28 Feb 2025 15:56:19 +0200 Subject: [PATCH] Fix something stupid and start implementing context schedule --- context.py | 184 ++++++++++++++++++++++++++++++++++++++ nodes.py | 126 ++++++++++++++++++++++---- wanvideo/modules/model.py | 14 +-- wanvideo/modules/vae.py | 7 +- 4 files changed, 300 insertions(+), 31 deletions(-) create mode 100644 context.py diff --git a/context.py b/context.py new file mode 100644 index 0000000..6a30fed --- /dev/null +++ b/context.py @@ -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)) + ) diff --git a/nodes.py b/nodes.py index b5c6519..c24ad83 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 6b65b48..ca9acc8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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) diff --git a/wanvideo/modules/vae.py b/wanvideo/modules/vae.py index 5c6da57..69f9bad 100644 --- a/wanvideo/modules/vae.py +++ b/wanvideo/modules/vae.py @@ -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)