Add enhance-a-video, start Spargeattn testing

This commit is contained in:
kijai
2025-02-28 00:49:17 +02:00
parent 9e957121be
commit dd3eedcd86
5 changed files with 197 additions and 10 deletions
View File
+53
View File
@@ -0,0 +1,53 @@
import torch
from einops import rearrange
from .globals import get_enhance_weight, get_num_frames
def get_feta_scores(query, key):
img_q, img_k = query, key
num_frames = get_num_frames()
B, S, N, C = img_q.shape
# Calculate spatial dimension
spatial_dim = S // num_frames
# Add time dimension between spatial and head dims
query_image = img_q.reshape(B, spatial_dim, num_frames, N, C)
key_image = img_k.reshape(B, spatial_dim, num_frames, N, C)
# Expand time dimension
query_image = query_image.expand(-1, -1, num_frames, -1, -1) # [B, S, T, N, C]
key_image = key_image.expand(-1, -1, num_frames, -1, -1) # [B, S, T, N, C]
# Reshape to match feta_score input format: [(B S) N T C]
query_image = rearrange(query_image, "b s t n c -> (b s) n t c") #torch.Size([3200, 24, 5, 128])
key_image = rearrange(key_image, "b s t n c -> (b s) n t c")
return feta_score(query_image, key_image, C, num_frames)
def feta_score(query_image, key_image, head_dim, num_frames):
scale = head_dim**-0.5
query_image = query_image * scale
attn_temp = query_image @ key_image.transpose(-2, -1) # translate attn to float32
attn_temp = attn_temp.to(torch.float32)
attn_temp = attn_temp.softmax(dim=-1)
# Reshape to [batch_size * num_tokens, num_frames, num_frames]
attn_temp = attn_temp.reshape(-1, num_frames, num_frames)
# Create a mask for diagonal elements
diag_mask = torch.eye(num_frames, device=attn_temp.device).bool()
diag_mask = diag_mask.unsqueeze(0).expand(attn_temp.shape[0], -1, -1)
# Zero out diagonal elements
attn_wo_diag = attn_temp.masked_fill(diag_mask, 0)
# Calculate mean for each token's attention matrix
# Number of off-diagonal elements per matrix is n*n - n
num_off_diag = num_frames * num_frames - num_frames
mean_scores = attn_wo_diag.sum(dim=(1, 2)) / num_off_diag
enhance_scores = mean_scores.mean() * (num_frames + get_enhance_weight())
enhance_scores = enhance_scores.clamp(min=1)
return enhance_scores
+36
View File
@@ -0,0 +1,36 @@
import torch
NUM_FRAMES = None
FETA_WEIGHT = None
ENABLE_FETA= False
@torch.compiler.disable()
def set_num_frames(num_frames: int):
global NUM_FRAMES
NUM_FRAMES = num_frames
@torch.compiler.disable()
def get_num_frames() -> int:
return NUM_FRAMES
def enable_enhance():
global ENABLE_FETA
ENABLE_FETA = True
def disable_enhance():
global ENABLE_FETA
ENABLE_FETA = False
@torch.compiler.disable()
def is_enhance_enabled() -> bool:
return ENABLE_FETA
@torch.compiler.disable()
def set_enhance_weight(feta_weight: float):
global FETA_WEIGHT
FETA_WEIGHT = feta_weight
@torch.compiler.disable()
def get_enhance_weight() -> float:
return FETA_WEIGHT
+71 -3
View File
@@ -13,6 +13,8 @@ from .wanvideo.utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
get_sampling_sigmas, retrieve_timesteps)
from .wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight, set_num_frames
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
@@ -153,6 +155,24 @@ def standardize_lora_key_format(lora_sd):
new_sd[k] = v
return new_sd
class WanVideoEnhanceAVideo:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"weight": ("FLOAT", {"default": 2.0, "min": 0, "max": 100, "step": 0.01, "tooltip": "The feta Weight of the Enhance-A-Video"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply Enhance-A-Video"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply Enhance-A-Video"}),
},
}
RETURN_TYPES = ("FETAARGS",)
RETURN_NAMES = ("feta_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/NUS-HPC-AI-Lab/Enhance-A-Video"
def setargs(self, **kwargs):
return (kwargs, )
class WanVideoLoraSelect:
@classmethod
@@ -234,6 +254,8 @@ class WanVideoModelLoader:
"flash_attn_2",
"flash_attn_3",
"sageattn",
"spargeattn",
"spargeattn_tune",
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
@@ -833,6 +855,7 @@ class WanVideoSampler:
"optional": {
"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", ),
}
}
@@ -841,7 +864,8 @@ class WanVideoSampler:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, force_offload=True, samples=None, denoise_strength=1.0):
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):
patcher = model
model = model.model
transformer = model.diffusion_model
@@ -879,7 +903,7 @@ class WanVideoSampler:
if denoise_strength < 1.0:
steps = int(steps * denoise_strength)
timesteps = timesteps[-(steps + 1):]
timesteps = timesteps[-(steps + 1):]
seed_g = torch.Generator(device=torch.device("cpu"))
seed_g.manual_seed(seed)
@@ -941,7 +965,6 @@ class WanVideoSampler:
arg_c = base_args.copy()
arg_c.update({'context': [text_embeds["prompt_embeds"][0]]})
arg_null = base_args.copy()
arg_null.update({'context': text_embeds["negative_prompt_embeds"]})
@@ -962,9 +985,37 @@ class WanVideoSampler:
else:
if model["manual_offloading"]:
transformer.to(device)
#feta
if feta_args is not None:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
feta_end_percent = feta_args["end_percent"]
set_num_frames(latent.shape[2])
enable_enhance()
else:
disable_enhance()
mm.soft_empty_cache()
gc.collect()
if "sparge" in transformer.attention_mode:
from spas_sage_attn.autotune import (
SparseAttentionMeansim,
extract_sparse_attention_state_dict,
load_sparse_attention_state_dict,
)
for idx, block in enumerate(transformer.blocks):
block.self_attn.verbose = True
block.self_attn.inner_attention = SparseAttentionMeansim(l1=0.06, pv_l1=0.065)
if transformer.attention_mode == "spargeattn":
saved_state_dict = torch.load("sparge_wan_30_steps_1_iter.pt")
for key in saved_state_dict.keys():
print(key)
load_sparse_attention_state_dict(transformer, saved_state_dict, verbose = True)
#for idx, block in enumerate(transformer.blocks):
# print(f"Block {idx} attn1: {block}")
try:
torch.cuda.reset_peak_memory_stats(device)
@@ -977,7 +1028,15 @@ class WanVideoSampler:
timestep = [t]
timestep = torch.stack(timestep).to(device)
current_step_percentage = i / len(timesteps)
if feta_args is not None:
if feta_start_percent <= current_step_percentage <= feta_end_percent:
enable_enhance()
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:
@@ -988,6 +1047,7 @@ class WanVideoSampler:
noise_pred_cond - noise_pred_uncond)
else:
noise_pred = noise_pred_cond
#model inference end
latent = latent.to(offload_device)
@@ -1008,6 +1068,11 @@ class WanVideoSampler:
pbar.update(1)
del latent_model_input, timestep
if transformer.attention_mode == "spargeattn_tune":
saved_state_dict = extract_sparse_attention_state_dict(transformer)
torch.save(saved_state_dict, "sparge_wan.pt")
save_torch_file(saved_state_dict, "sparge_wan.safetensors")
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
@@ -1210,6 +1275,8 @@ NODE_CLASS_MAPPINGS = {
"WanVideoEmptyEmbeds": WanVideoEmptyEmbeds,
"WanVideoLoraSelect": WanVideoLoraSelect,
"WanVideoLoraBlockEdit": WanVideoLoraBlockEdit,
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -1228,4 +1295,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEmptyEmbeds": "WanVideo Empty Embeds",
"WanVideoLoraSelect": "WanVideo Lora Select",
"WanVideoLoraBlockEdit": "WanVideo Lora Block Edit",
"WanVideoEnhanceAVideo": "WanVideo Enhance AVideo"
}
+37 -7
View File
@@ -7,6 +7,9 @@ import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled, set_num_frames
from .attention import attention
__all__ = ['WanModel']
@@ -149,13 +152,40 @@ class WanSelfAttention(nn.Module):
q, k, v = qkv_fn(x)
x = attention(
q=rope_apply(q, grid_sizes, freqs),
k=rope_apply(k, grid_sizes, freqs),
v=v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode)
if is_enhance_enabled():
feta_scores = get_feta_scores(q, k)
if self.attention_mode == 'spargeattn_tune' or self.attention_mode == 'spargeattn':
tune_mode = False
if self.attention_mode == 'spargeattn_tune':
tune_mode = True
if hasattr(self, 'inner_attention'):
#print("has inner attention")
q=rope_apply(q, grid_sizes, freqs)
k=rope_apply(k, grid_sizes, freqs)
q = q.permute(0, 2, 1, 3)
k = k.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3)
x = self.inner_attention(
q=q,
k=k,
v=v,
is_causal=False,
tune_mode=tune_mode
).permute(0, 2, 1, 3)
#print("inner attention", x.shape) #inner attention torch.Size([1, 12, 32760, 128])
else:
x = attention(
q=rope_apply(q, grid_sizes, freqs),
k=rope_apply(k, grid_sizes, freqs),
v=v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode)
if is_enhance_enabled():
x *= feta_scores
# output
x = x.flatten(2)