Allow splitting self attention as well when using multiple prompts
This commit is contained in:
@@ -1592,6 +1592,24 @@ class WanVideoLoopArgs:
|
||||
|
||||
def process(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class WanVideoExperimentalArgs:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"video_attention_split_steps": ("STRING", {"default": "2, 3", "tooltip": "Steps to split self attention when using multiple prompts"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("EXPERIMENTALARGS", )
|
||||
RETURN_NAMES = ("exp_args",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Experimental stuff"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def process(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class WanVideoSampler:
|
||||
@classmethod
|
||||
@@ -1625,6 +1643,7 @@ class WanVideoSampler:
|
||||
"slg_args": ("SLGARGS", ),
|
||||
"rope_function": (["default", "comfy"], {"default": "default", "tooltip": "!EXPERIMENTAL! Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
|
||||
"loop_args": ("LOOPARGS", ),
|
||||
"experimental_args": ("EXPERIMENTALARGS", ),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1635,7 +1654,7 @@ class WanVideoSampler:
|
||||
|
||||
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, context_options=None,
|
||||
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None):
|
||||
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, experimental_args=None):
|
||||
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
|
||||
patcher = model
|
||||
model = model.model
|
||||
@@ -1956,6 +1975,11 @@ class WanVideoSampler:
|
||||
drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device)
|
||||
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
|
||||
|
||||
if experimental_args is not None:
|
||||
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
|
||||
if video_attention_split_steps:
|
||||
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
|
||||
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, teacache_state=None):
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True):
|
||||
nonlocal patcher
|
||||
@@ -2628,6 +2652,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoLoopArgs": WanVideoLoopArgs,
|
||||
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
|
||||
"WanVideoSetBlockSwap": WanVideoSetBlockSwap,
|
||||
"WanVideoExperimentalArgs": WanVideoExperimentalArgs,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoSampler": "WanVideo Sampler",
|
||||
@@ -2660,4 +2685,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLoopArgs": "WanVideo Loop Args",
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
"WanVideoSetBlockSwap": "WanVideo Set BlockSwap",
|
||||
"WanVideoExperimentalArgs": "WanVideo Experimental Args",
|
||||
}
|
||||
|
||||
+107
-38
@@ -221,7 +221,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, seq_chunks=1,current_step=0, video_attention_split_steps = [], rope_func = "default"):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
@@ -240,35 +240,84 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
if self.attention_mode == 'spargeattn_tune' or self.attention_mode == 'spargeattn':
|
||||
tune_mode = False
|
||||
if self.attention_mode == 'spargeattn_tune':
|
||||
tune_mode = True
|
||||
# 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])
|
||||
# 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:
|
||||
if rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
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)
|
||||
if is_enhance_enabled():
|
||||
feta_scores = get_feta_scores(q, k)
|
||||
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)
|
||||
# Split by frames
|
||||
if seq_chunks > 1 and current_step in video_attention_split_steps:
|
||||
outputs = []
|
||||
# Extract frame, height, width from grid_sizes - force to CPU scalars
|
||||
frames = grid_sizes[0][0].item()
|
||||
height = grid_sizes[0][1].item()
|
||||
width = grid_sizes[0][2].item()
|
||||
tokens_per_frame = height * width
|
||||
|
||||
actual_chunks = min(seq_chunks, frames)
|
||||
if isinstance(actual_chunks, torch.Tensor):
|
||||
actual_chunks = actual_chunks.item()
|
||||
|
||||
frame_chunks = [] # Pre-calculate all chunk boundaries
|
||||
start_frame = 0
|
||||
base_frames_per_chunk = frames // actual_chunks
|
||||
extra_frames = frames % actual_chunks
|
||||
|
||||
# Pre-calculate all chunks
|
||||
for i in range(actual_chunks):
|
||||
chunk_size = base_frames_per_chunk + (1 if i < extra_frames else 0)
|
||||
end_frame = start_frame + chunk_size
|
||||
frame_chunks.append((start_frame, end_frame))
|
||||
start_frame = end_frame
|
||||
|
||||
# Process each chunk using the pre-calculated boundaries
|
||||
for start_frame, end_frame in frame_chunks:
|
||||
# Convert to token indices
|
||||
start_idx = int(start_frame * tokens_per_frame)
|
||||
end_idx = int(end_frame * tokens_per_frame)
|
||||
|
||||
chunk_q = q[:, start_idx:end_idx, :, :]
|
||||
chunk_k = k[:, start_idx:end_idx, :, :]
|
||||
chunk_v = v[:, start_idx:end_idx, :, :]
|
||||
|
||||
chunk_out = attention(
|
||||
q=chunk_q,
|
||||
k=chunk_k,
|
||||
v=chunk_v,
|
||||
k_lens=seq_lens,
|
||||
window_size=self.window_size,
|
||||
attention_mode=self.attention_mode)
|
||||
|
||||
outputs.append(chunk_out)
|
||||
|
||||
# Concatenate outputs along the sequence dimension
|
||||
x = torch.cat(outputs, dim=1)
|
||||
else:
|
||||
# Original attention computation
|
||||
x = attention(
|
||||
q=q,
|
||||
k=k,
|
||||
@@ -416,8 +465,11 @@ class WanAttentionBlock(nn.Module):
|
||||
freqs,
|
||||
context,
|
||||
context_lens,
|
||||
current_step,
|
||||
video_attention_split_steps=[],
|
||||
rope_func = "default",
|
||||
clip_embed=None,
|
||||
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -433,29 +485,42 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
# self-attention
|
||||
y = self.self_attn(
|
||||
self.norm1(x).float() * (1 + e[1]) + e[0], seq_lens, grid_sizes,
|
||||
freqs, rope_func=rope_func)
|
||||
self.norm1(x).float() * (1 + e[1]) + e[0],
|
||||
seq_lens, grid_sizes,
|
||||
freqs, rope_func=rope_func,
|
||||
seq_chunks=max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0),
|
||||
current_step=current_step,
|
||||
video_attention_split_steps=video_attention_split_steps
|
||||
)
|
||||
|
||||
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, clip_embed=None):
|
||||
def cross_attn_ffn(x, context, context_lens, e, clip_embed=None, grid_sizes=None):
|
||||
if context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1):
|
||||
# Get number of prompts
|
||||
num_prompts = context.shape[0]
|
||||
num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0]
|
||||
num_segments = max(num_prompts, num_clip_embeds)
|
||||
|
||||
# split the sequence dimension
|
||||
seq_len = x.shape[1]
|
||||
segment_length = seq_len // num_prompts
|
||||
# Extract spatial dimensions
|
||||
frames, height, width = grid_sizes[0] # Assuming batch size 1
|
||||
tokens_per_frame = height * width
|
||||
|
||||
# Distribute frames across prompts
|
||||
frames_per_segment = max(1, frames // num_segments)
|
||||
|
||||
# Process each prompt segment
|
||||
x_combined = torch.zeros_like(x)
|
||||
|
||||
for i in range(num_segments):
|
||||
# Calculate indices for this segment
|
||||
start_idx = i * segment_length
|
||||
end_idx = (i+1) * segment_length if i < num_segments-1 else seq_len
|
||||
# Calculate frame boundaries for this segment
|
||||
start_frame = i * frames_per_segment
|
||||
end_frame = min((i+1) * frames_per_segment, frames) if i < num_segments-1 else frames
|
||||
|
||||
# Convert frame indices to token indices
|
||||
start_idx = start_frame * tokens_per_frame
|
||||
end_idx = end_frame * tokens_per_frame
|
||||
segment_indices = torch.arange(start_idx, end_idx, device=x.device, dtype=torch.long)
|
||||
|
||||
# Get prompt segment (cycle through available prompts if needed)
|
||||
@@ -494,7 +559,7 @@ class WanAttentionBlock(nn.Module):
|
||||
x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32))
|
||||
return x
|
||||
|
||||
x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed)
|
||||
x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes)
|
||||
return x
|
||||
|
||||
|
||||
@@ -654,6 +719,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
|
||||
self.use_non_blocking = True
|
||||
|
||||
self.video_attention_split_steps = []
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
@@ -900,7 +967,9 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
context=context,
|
||||
context_lens=context_lens,
|
||||
clip_embed=clip_embed,
|
||||
rope_func=rope_func
|
||||
rope_func=rope_func,
|
||||
current_step=current_step,
|
||||
video_attention_split_steps=self.video_attention_split_steps,
|
||||
)
|
||||
|
||||
for b, block in enumerate(self.blocks):
|
||||
|
||||
Reference in New Issue
Block a user