Add flex_attention for CausVid

This commit is contained in:
kijai
2025-05-13 21:35:34 +03:00
parent f2bc29b931
commit bfed2afc2e
3 changed files with 208 additions and 20 deletions
+11 -5
View File
@@ -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
View File
@@ -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:
+84
View File
@@ -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