Split cross attention when using multiple prompts and/or clip image embeds
Allows for (pretty weak) prompt travel and may help with the endframe workflow.
This commit is contained in:
@@ -894,7 +894,7 @@ class WanVideoTextEncode:
|
||||
RETURN_NAMES = ("text_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Encodes text prompts into text embeddings. For context windowing you can input multiple prompts separated by '|'"
|
||||
DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length"
|
||||
|
||||
def process(self, t5, positive_prompt, negative_prompt,force_offload=True, model_to_offload=None):
|
||||
|
||||
@@ -918,10 +918,6 @@ class WanVideoTextEncode:
|
||||
context = encoder(positive_prompts, device)
|
||||
context_null = encoder([negative_prompt], device)
|
||||
|
||||
|
||||
context = [t.to(device) for t in context]
|
||||
context_null = [t.to(device) for t in context_null]
|
||||
|
||||
if force_offload:
|
||||
encoder.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
@@ -1141,11 +1137,12 @@ class WanVideoClipVisionEncode:
|
||||
"strength_1": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional clip embed multiplier"}),
|
||||
"strength_2": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional clip embed multiplier"}),
|
||||
"crop": (["center", "disabled"], {"default": "center", "tooltip": "Crop image to 224x224 before encoding"}),
|
||||
"combine_embeds": (["average", "sum", "concat"], {"default": "average", "tooltip": "Method to combine multiple clip embeds"}),
|
||||
"combine_embeds": (["average", "sum", "concat", "batch"], {"default": "average", "tooltip": "Method to combine multiple clip embeds"}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"image_2": ("IMAGE", ),
|
||||
"negative_image": ("IMAGE", {"tooltip": "image to use for uncond"}),
|
||||
"tiles": ("INT", {"default": 0, "min": 0, "max": 16, "step": 2, "tooltip": "Use matteo's tiled image encoding for improved accuracy"}),
|
||||
"ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Ratio of the tile average"}),
|
||||
}
|
||||
@@ -1156,7 +1153,7 @@ class WanVideoClipVisionEncode:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, clip_vision, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None, tiles=0, ratio=1.0):
|
||||
def process(self, clip_vision, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None, negative_image=None, tiles=0, ratio=1.0):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
@@ -1170,18 +1167,26 @@ class WanVideoClipVisionEncode:
|
||||
image = image_1
|
||||
|
||||
clip_vision.model.to(device)
|
||||
image = image.to(device)
|
||||
negative_clip_embeds = None
|
||||
|
||||
if tiles > 0:
|
||||
log.info("Using tiled image encoding")
|
||||
clip_embeds = clip_encode_image_tiled(clip_vision, image, tiles=tiles, ratio=ratio)
|
||||
clip_embeds = clip_encode_image_tiled(clip_vision, image.to(device), tiles=tiles, ratio=ratio)
|
||||
if negative_image is not None:
|
||||
negative_clip_embeds = clip_encode_image_tiled(clip_vision, negative_image.to(device), tiles=tiles, ratio=ratio)
|
||||
else:
|
||||
if isinstance(clip_vision, ClipVisionModel):
|
||||
clip_embeds = clip_vision.encode_image(image).last_hidden_state.to(device)
|
||||
if negative_image is not None:
|
||||
negative_clip_embeds = clip_vision.encode_image(negative_image).last_hidden_state.to(device)
|
||||
else:
|
||||
pixel_values = clip_preprocess(image.to(device), size=224, mean=image_mean, std=image_std, crop=(not crop == "disabled")).float()
|
||||
clip_embeds = clip_vision.visual(pixel_values)
|
||||
if negative_image is not None:
|
||||
pixel_values = clip_preprocess(negative_image.to(device), size=224, mean=image_mean, std=image_std, crop=(not crop == "disabled")).float()
|
||||
negative_clip_embeds = clip_vision.visual(pixel_values)
|
||||
log.info(f"Clip embeds shape: {clip_embeds.shape}")
|
||||
|
||||
|
||||
if clip_embeds.shape[0] > 1:
|
||||
embed_1 = clip_embeds[0:1] * strength_1
|
||||
@@ -1192,6 +1197,9 @@ class WanVideoClipVisionEncode:
|
||||
clip_embeds = torch.sum(torch.stack([embed_1, embed_2]), dim=0)
|
||||
elif combine_embeds == "concat":
|
||||
clip_embeds = torch.cat([embed_1, embed_2], dim=1)
|
||||
elif combine_embeds == "batch":
|
||||
clip_embeds = torch.cat([embed_1, embed_2], dim=0)
|
||||
|
||||
|
||||
log.info(f"Combined clip embeds shape: {clip_embeds.shape}")
|
||||
|
||||
@@ -1199,7 +1207,12 @@ class WanVideoClipVisionEncode:
|
||||
clip_vision.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
return (clip_embeds,)
|
||||
clip_embeds_dict = {
|
||||
"clip_embeds": clip_embeds,
|
||||
"negative_clip_embeds": negative_clip_embeds
|
||||
}
|
||||
|
||||
return (clip_embeds_dict,)
|
||||
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
@@ -1294,7 +1307,8 @@ class WanVideoImageToVideoEncode:
|
||||
|
||||
image_embeds = {
|
||||
"image_embeds": y,
|
||||
"clip_context": clip_embeds,
|
||||
"clip_context": clip_embeds.get("clip_embeds", None),
|
||||
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None),
|
||||
"max_seq_len": max_seq_len,
|
||||
"num_frames": num_frames,
|
||||
"lat_h": lat_h,
|
||||
@@ -1581,10 +1595,9 @@ class WanVideoSampler:
|
||||
|
||||
seed_g = torch.Generator(device=torch.device("cpu"))
|
||||
seed_g.manual_seed(seed)
|
||||
image_cond = None
|
||||
clip_fea = None
|
||||
control_latents = None
|
||||
end_image = None
|
||||
|
||||
image_cond, control_latents, clip_fea, clip_fea_neg, end_image = None, None, None, None, None
|
||||
|
||||
if transformer.model_type == "i2v":
|
||||
end_image = image_embeds.get("end_image", None)
|
||||
lat_h = image_embeds.get("lat_h", None)
|
||||
@@ -1603,7 +1616,7 @@ class WanVideoSampler:
|
||||
image_cond = image_embeds.get("image_embeds", None)
|
||||
print("image_cond", image_cond.shape)
|
||||
clip_fea = image_embeds.get("clip_context", None)
|
||||
|
||||
clip_fea_neg = image_embeds.get("negative_clip_context", None)
|
||||
else: #t2v
|
||||
target_shape = image_embeds.get("target_shape", None)
|
||||
if target_shape is None:
|
||||
@@ -1820,7 +1833,7 @@ class WanVideoSampler:
|
||||
source_embeds = flowedit_args["source_embeds"]
|
||||
source_image_embeds = flowedit_args.get("source_image_embeds", image_embeds)
|
||||
source_image_cond = source_image_embeds.get("image_embeds", None)
|
||||
source_clip_fea = source_image_embeds.get("clip_fea", None)
|
||||
source_clip_fea = source_image_embeds.get("clip_fea", clip_fea)
|
||||
skip_steps = flowedit_args["skip_steps"]
|
||||
drift_steps = flowedit_args["drift_steps"]
|
||||
source_cfg = flowedit_args["source_cfg"]
|
||||
@@ -1870,7 +1883,6 @@ class WanVideoSampler:
|
||||
patcher.model.is_patched = True
|
||||
|
||||
base_params = {
|
||||
'clip_fea': clip_fea,
|
||||
'seq_len': seq_len,
|
||||
'device': device,
|
||||
'freqs': freqs,
|
||||
@@ -1883,7 +1895,7 @@ class WanVideoSampler:
|
||||
if not batched_cfg:
|
||||
#cond
|
||||
noise_pred_cond, teacache_state_cond = transformer(
|
||||
[z], context=[positive_embeds], is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
[z], context=positive_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=teacache_state[0] if teacache_state else None,
|
||||
**base_params
|
||||
)
|
||||
@@ -1892,7 +1904,8 @@ class WanVideoSampler:
|
||||
return noise_pred_cond, [teacache_state_cond]
|
||||
#uncond
|
||||
noise_pred_uncond, teacache_state_uncond = transformer(
|
||||
[z], context=negative_embeds, is_uncond=True, current_step_percentage=current_step_percentage,
|
||||
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
is_uncond=True, current_step_percentage=current_step_percentage,
|
||||
pred_id=teacache_state[1] if teacache_state else None,
|
||||
**base_params
|
||||
)
|
||||
@@ -1901,7 +1914,7 @@ class WanVideoSampler:
|
||||
#batched
|
||||
else:
|
||||
[noise_pred_cond, noise_pred_uncond], teacache_state_cond = transformer(
|
||||
[z] + [z], context= [positive_embeds] + negative_embeds, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
[z] + [z], context= positive_embeds + negative_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=teacache_state[0] if teacache_state else None,
|
||||
**base_params
|
||||
)
|
||||
@@ -1999,7 +2012,10 @@ class WanVideoSampler:
|
||||
if context_options["verbose"]:
|
||||
log.info(f"Prompt index: {prompt_index}")
|
||||
|
||||
positive = source_embeds["prompt_embeds"][prompt_index]
|
||||
if len(source_embeds["prompt_embeds"]) > 1:
|
||||
positive = source_embeds["prompt_embeds"][prompt_index]
|
||||
else:
|
||||
positive = source_embeds["prompt_embeds"]
|
||||
|
||||
partial_img_emb = None
|
||||
if source_image_cond is not None:
|
||||
@@ -2023,7 +2039,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
vt_src, self.teacache_state_source = predict_with_cfg(
|
||||
zt_src, cfg[idx],
|
||||
source_embeds["prompt_embeds"][0],
|
||||
source_embeds["prompt_embeds"],
|
||||
source_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, source_image_cond,
|
||||
source_clip_fea,
|
||||
@@ -2050,7 +2066,10 @@ class WanVideoSampler:
|
||||
if context_options["verbose"]:
|
||||
log.info(f"Prompt index: {prompt_index}")
|
||||
|
||||
positive = text_embeds["prompt_embeds"][prompt_index]
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
positive = text_embeds["prompt_embeds"][prompt_index]
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
partial_img_emb = None
|
||||
if image_cond is not None:
|
||||
@@ -2074,7 +2093,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
vt_tgt, self.teacache_state = predict_with_cfg(
|
||||
zt_tgt, cfg[idx],
|
||||
text_embeds["prompt_embeds"][0],
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea,
|
||||
teacache_state=self.teacache_state)
|
||||
@@ -2102,7 +2121,10 @@ class WanVideoSampler:
|
||||
log.info(f"Prompt index: {prompt_index}")
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
positive = text_embeds["prompt_embeds"][prompt_index]
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
positive = text_embeds["prompt_embeds"][prompt_index]
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
partial_img_emb = None
|
||||
if image_cond is not None:
|
||||
@@ -2166,7 +2188,7 @@ class WanVideoSampler:
|
||||
noise_pred, self.teacache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"][0],
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea,
|
||||
teacache_state=self.teacache_state)
|
||||
|
||||
+67
-18
@@ -289,7 +289,7 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, context_lens, clip_fea_tokens=None):
|
||||
def forward(self, x, context, context_lens, clip_embed=None):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -329,23 +329,23 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
def forward(self, x, context, context_lens, clip_fea_tokens=257):
|
||||
def forward(self, x, context, context_lens, clip_embed):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
"""
|
||||
context_img = context[:, :clip_fea_tokens]
|
||||
context = context[:, clip_fea_tokens:]
|
||||
#context_img = context[:, :clip_embed.shape[1]]
|
||||
#context = context[:, clip_embed.shape[1]:]
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
k_img = self.norm_k_img(self.k_img(context_img)).view(b, -1, n, d)
|
||||
v_img = self.v_img(context_img).view(b, -1, n, d)
|
||||
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
|
||||
v_img = self.v_img(clip_embed).view(b, -1, n, d)
|
||||
img_x = attention(q, k_img, v_img, k_lens=None, attention_mode=self.attention_mode)
|
||||
# compute attention
|
||||
x = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode)
|
||||
@@ -417,7 +417,7 @@ class WanAttentionBlock(nn.Module):
|
||||
context,
|
||||
context_lens,
|
||||
rope_func = "default",
|
||||
clip_fea_tokens=257,
|
||||
clip_embed=None,
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -438,13 +438,63 @@ class WanAttentionBlock(nn.Module):
|
||||
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_fea_tokens=clip_fea_tokens):
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_fea_tokens=clip_fea_tokens)
|
||||
y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3])
|
||||
x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32))
|
||||
return x
|
||||
def cross_attn_ffn(x, context, context_lens, e, clip_embed=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
|
||||
|
||||
# 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
|
||||
segment_indices = torch.arange(start_idx, end_idx, device=x.device, dtype=torch.long)
|
||||
|
||||
# Get prompt segment (cycle through available prompts if needed)
|
||||
prompt_idx = i % num_prompts
|
||||
segment_context = context[prompt_idx:prompt_idx+1]
|
||||
segment_context_lens = None
|
||||
if context_lens is not None:
|
||||
segment_context_lens = context_lens[prompt_idx:prompt_idx+1]
|
||||
|
||||
# Handle clip_embed for this segment (cycle through available embeddings)
|
||||
segment_clip_embed = None
|
||||
if clip_embed is not None:
|
||||
clip_idx = i % num_clip_embeds
|
||||
segment_clip_embed = clip_embed[clip_idx:clip_idx+1]
|
||||
|
||||
# Get tensor segment
|
||||
x_segment = x[:, segment_indices, :]
|
||||
|
||||
# Process segment with its prompt and clip embedding
|
||||
processed_segment = self.cross_attn(self.norm3(x_segment), segment_context, segment_context_lens, clip_embed=segment_clip_embed)
|
||||
processed_segment = processed_segment.to(x.dtype)
|
||||
|
||||
# Add to combined result
|
||||
x_combined[:, segment_indices, :] = processed_segment
|
||||
|
||||
# Continue with FFN
|
||||
x = x + x_combined
|
||||
y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3])
|
||||
x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32))
|
||||
return x
|
||||
|
||||
else:
|
||||
cross_attn_result = self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed)
|
||||
x = x + cross_attn_result
|
||||
y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3])
|
||||
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_fea_tokens=clip_fea_tokens)
|
||||
x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed)
|
||||
return x
|
||||
|
||||
|
||||
@@ -789,14 +839,13 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
if self.offload_txt_emb:
|
||||
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
clip_fea_tokens = 257
|
||||
clip_embed = None
|
||||
if clip_fea is not None:
|
||||
clip_fea_tokens = clip_fea.shape[1]
|
||||
clip_fea = clip_fea.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
context = torch.concat([context_clip, context], dim=1)
|
||||
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
#context = torch.concat([context_clip, context], dim=1)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
@@ -850,7 +899,7 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
freqs=freqs,
|
||||
context=context,
|
||||
context_lens=context_lens,
|
||||
clip_fea_tokens=clip_fea_tokens,
|
||||
clip_embed=clip_embed,
|
||||
rope_func=rope_func
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user