Add flex_attention for CausVid
This commit is contained in:
@@ -13,6 +13,7 @@ from .wanvideo.modules.t5 import T5EncoderModel
|
||||
from .wanvideo.utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
|
||||
get_sampling_sigmas, retrieve_timesteps)
|
||||
from .wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
from .wanvideo.utils.basic_flowmatch import FlowMatchScheduler
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler
|
||||
from .wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
|
||||
|
||||
@@ -474,6 +475,7 @@ class WanVideoModelLoader:
|
||||
"flash_attn_2",
|
||||
"flash_attn_3",
|
||||
"sageattn",
|
||||
"flex_attention",
|
||||
#"spargeattn", needs tuning
|
||||
#"spargeattn_tune",
|
||||
], {"default": "sdpa"}),
|
||||
@@ -2309,7 +2311,7 @@ class WanVideoSampler:
|
||||
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
|
||||
"scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "deis", "lcm", "lcm/beta"],
|
||||
"scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "deis", "lcm", "lcm/beta", "flowmatch_causvid"],
|
||||
{
|
||||
"default": 'unipc'
|
||||
}),
|
||||
@@ -2390,10 +2392,14 @@ class WanVideoSampler:
|
||||
sample_scheduler.sigmas[-1] = 1e-6
|
||||
elif 'lcm' in scheduler:
|
||||
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
elif 'flowmatch_causvid' in scheduler:
|
||||
sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
|
||||
sample_scheduler.timesteps = torch.tensor([1000, 934, 862, 756, 603, 410, 250, 140, 74, 0])[:steps].to(device)
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
print("timesteps: ", timesteps)
|
||||
|
||||
if denoise_strength < 1.0:
|
||||
steps = int(steps * denoise_strength)
|
||||
@@ -3297,13 +3303,13 @@ class WanVideoSampler:
|
||||
step_args = {
|
||||
"generator": seed_g,
|
||||
}
|
||||
if isinstance(sample_scheduler, DEISMultistepScheduler):
|
||||
if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler):
|
||||
step_args.pop("generator", None)
|
||||
temp_x0 = sample_scheduler.step(
|
||||
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0),
|
||||
t,
|
||||
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0),
|
||||
return_dict=False,
|
||||
#return_dict=False,
|
||||
**step_args)[0]
|
||||
latent = temp_x0.squeeze(0)
|
||||
|
||||
@@ -3312,7 +3318,7 @@ class WanVideoSampler:
|
||||
if recammaster is not None:
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len] - noise_pred[:, :orig_noise_len].to(t.device) * t / 1000).detach().permute(1,0,2,3)
|
||||
else:
|
||||
callback_latent = (latent_model_input - noise_pred.to(t.device) * t / 1000).detach().permute(1,0,2,3)
|
||||
callback_latent = (latent_model_input.cpu() - noise_pred.cpu() * t.cpu() / 1000).detach().permute(1,0,2,3)
|
||||
callback(idx, callback_latent, None, steps)
|
||||
else:
|
||||
pbar.update(1)
|
||||
|
||||
+113
-15
@@ -9,6 +9,13 @@ from einops import repeat, rearrange
|
||||
from ...enhance_a_video.enhance import get_feta_scores
|
||||
from ...enhance_a_video.globals import is_enhance_enabled
|
||||
|
||||
try:
|
||||
from torch.nn.attention.flex_attention import create_block_mask, flex_attention, BlockMask
|
||||
create_block_mask = torch.compile(create_block_mask)
|
||||
flex_attention = torch.compile(flex_attention)
|
||||
except:
|
||||
pass
|
||||
|
||||
from .attention import attention
|
||||
import numpy as np
|
||||
__all__ = ['WanModel']
|
||||
@@ -180,7 +187,7 @@ class WanSelfAttention(nn.Module):
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default"):
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default", block_mask=None):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
@@ -199,22 +206,58 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
if rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
else:
|
||||
q=rope_apply(q, grid_sizes, freqs)
|
||||
k=rope_apply(k, grid_sizes, freqs)
|
||||
|
||||
if is_enhance_enabled():
|
||||
feta_scores = get_feta_scores(q, k)
|
||||
|
||||
x = attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
k_lens=seq_lens,
|
||||
window_size=self.window_size,
|
||||
attention_mode=self.attention_mode)
|
||||
if self.attention_mode == 'flex_attention':
|
||||
if rope_func == "comfy":
|
||||
roped_query, roped_key = apply_rope_comfy(q, k, freqs)
|
||||
else:
|
||||
roped_query = rope_apply(q, grid_sizes, freqs).type_as(v)
|
||||
roped_key = rope_apply(k, grid_sizes, freqs).type_as(v)
|
||||
|
||||
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
|
||||
padded_roped_query = torch.cat(
|
||||
[roped_query,
|
||||
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
|
||||
device=q.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
padded_roped_key = torch.cat(
|
||||
[roped_key, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
|
||||
device=k.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
padded_v = torch.cat(
|
||||
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
|
||||
device=v.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
x = flex_attention(
|
||||
query=padded_roped_query.transpose(2, 1),
|
||||
key=padded_roped_key.transpose(2, 1),
|
||||
value=padded_v.transpose(2, 1),
|
||||
block_mask=block_mask
|
||||
)[:, :, :-padded_length].transpose(2, 1)
|
||||
|
||||
else:
|
||||
if rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
else:
|
||||
q=rope_apply(q, grid_sizes, freqs)
|
||||
k=rope_apply(k, grid_sizes, freqs)
|
||||
|
||||
x = attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
k_lens=seq_lens,
|
||||
window_size=self.window_size,
|
||||
attention_mode=self.attention_mode)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -512,6 +555,7 @@ class WanAttentionBlock(nn.Module):
|
||||
audio_context_lens=None,
|
||||
audio_scale=1.0,
|
||||
num_latent_frames=21,
|
||||
block_mask=None
|
||||
|
||||
):
|
||||
r"""
|
||||
@@ -549,7 +593,8 @@ class WanAttentionBlock(nn.Module):
|
||||
y = self.self_attn.forward(
|
||||
input_x,
|
||||
seq_lens, grid_sizes,
|
||||
freqs, rope_func=rope_func
|
||||
freqs, rope_func=rope_func,
|
||||
block_mask=block_mask
|
||||
)
|
||||
#ReCamMaster
|
||||
if camera_embed is not None:
|
||||
@@ -964,6 +1009,50 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
else:
|
||||
self.control_adapter = None
|
||||
|
||||
self.block_mask=None
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
frame_seqlen: int = 1560, num_frame_per_block=1
|
||||
) -> BlockMask:
|
||||
"""
|
||||
we will divide the token sequence into the following format
|
||||
[1 latent frame] [1 latent frame] ... [1 latent frame]
|
||||
We use flexattention to construct the attention mask
|
||||
"""
|
||||
print("num_frames", num_frames)
|
||||
print("frame_seqlen", frame_seqlen)
|
||||
total_length = num_frames * frame_seqlen
|
||||
|
||||
# we do right padding to get to a multiple of 128
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
|
||||
ends = torch.zeros(total_length + padded_length,
|
||||
device=device, dtype=torch.long)
|
||||
|
||||
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
|
||||
frame_indices = torch.arange(
|
||||
start=0,
|
||||
end=total_length,
|
||||
step=frame_seqlen * num_frame_per_block,
|
||||
device=device
|
||||
)
|
||||
|
||||
for tmp in frame_indices:
|
||||
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
|
||||
frame_seqlen * num_frame_per_block
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
|
||||
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
|
||||
|
||||
|
||||
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
|
||||
KV_LEN=total_length + padded_length, _compile=False, device=device)
|
||||
|
||||
return block_mask
|
||||
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None):
|
||||
log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
@@ -1100,6 +1189,14 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
freqs = freqs.to(device)
|
||||
|
||||
_, F, H, W = x[0].shape
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.attention_mode == 'flex_attention' and current_step == 0:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device, num_frames=F,
|
||||
frame_seqlen=H * W // (self.patch_size[1] * self.patch_size[2]),
|
||||
num_frame_per_block=3
|
||||
)
|
||||
|
||||
if y is not None:
|
||||
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
|
||||
@@ -1293,7 +1390,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
audio_proj=audio_proj,
|
||||
audio_context_lens=audio_context_lens,
|
||||
num_latent_frames = F,
|
||||
audio_scale=audio_scale
|
||||
audio_scale=audio_scale,
|
||||
block_mask=self.block_mask
|
||||
)
|
||||
|
||||
if vace_data is not None:
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
"""
|
||||
The following code is copied from https://github.com/modelscope/DiffSynth-Studio/blob/main/diffsynth/schedulers/flow_match.py
|
||||
"""
|
||||
import torch
|
||||
|
||||
|
||||
class FlowMatchScheduler():
|
||||
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.inverse_timesteps = inverse_timesteps
|
||||
self.extra_one_step = extra_one_step
|
||||
self.reverse_sigmas = reverse_sigmas
|
||||
self.set_timesteps(num_inference_steps)
|
||||
|
||||
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False):
|
||||
sigma_start = self.sigma_min + \
|
||||
(self.sigma_max - self.sigma_min) * denoising_strength
|
||||
if self.extra_one_step:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
|
||||
else:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
|
||||
num_inference_steps) ** 2)
|
||||
y_shifted = y - y.min()
|
||||
bsmntw_weighing = y_shifted * \
|
||||
(num_inference_steps / y_shifted.sum())
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
|
||||
def step(self, model_output, timestep, sample, to_final=False):
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps - timestep).abs(), dim=0)
|
||||
print("timestep_id: ", timestep_id)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
|
||||
sigma_ = 1 if (
|
||||
self.inverse_timesteps or self.reverse_sigmas) else 0
|
||||
else:
|
||||
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
prev_sample = sample + model_output * (sigma_ - sigma)
|
||||
return prev_sample
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B, C, H, W]
|
||||
- noise: the noise with shape [B, C, H, W]
|
||||
- timestep: the timestep with shape [B]
|
||||
Output: the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
def training_weight(self, timestep):
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps - timestep.to(self.timesteps.device)).abs())
|
||||
weights = self.linear_timesteps_weights[timestep_id]
|
||||
return weights
|
||||
Reference in New Issue
Block a user