Add enhance-a-video, start Spargeattn testing
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user