Fix something stupid and start implementing context schedule
This commit is contained in:
+184
@@ -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))
|
||||
)
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user