Support EchoShot

https://github.com/D2I-ai/EchoShot
This commit is contained in:
kijai
2025-07-26 18:51:29 +03:00
parent 9a1ab1c656
commit df1476bf42
3 changed files with 164 additions and 49 deletions
+65
View File
@@ -0,0 +1,65 @@
import torch
from comfy.model_management import get_autocast_device, get_torch_device
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_z(x, grid_sizes, freqs, inner_t, shift=6):
n, c = x.size(2), x.size(3) // 2
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)
)
start_ind = [sum(inner_t[i][:_]) for _ in range(len(inner_t[i]))]
end_ind = [sum(inner_t[i][:_+1]) for _ in range(len(inner_t[i]))]
freq_select = []
for shot_ind, (s, e) in enumerate(zip(start_ind, end_ind)):
freq_select += [shot_ind * shift] * (e - s)
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(f, 1, 1, -1).expand(f, h, w, -1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output).float()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_c(x, freqs, inner_c, shift=6):
b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2
# loop over samples
output = []
for i in range(b):
# precompute multipliers
x_i = torch.view_as_complex(
x[i].to(torch.float64).reshape(s, n, -1, 2)
)
freq_select = []
for shot_ind, c_len in enumerate(inner_c[i]):
freq_select += [shot_ind * shift] * c_len
freq_select += [shot_ind+10] * (s-len(freq_select)) # extra suppression for the empty token
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(s, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
# append to collection
output.append(x_i)
return torch.stack(output).float()
+31 -6
View File
@@ -231,20 +231,34 @@ class WanVideoTextEncode:
if model_to_offload is not None:
log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device)
mm.soft_empty_cache()
encoder = t5["model"]
dtype = t5["dtype"]
echoshot = False
# Split positive prompts and process each with weights
positive_prompts_raw = [p.strip() for p in positive_prompt.split('|')]
positive_prompts = []
all_weights = []
# Split positive prompts and process each with weights
if "|" in positive_prompt:
log.info("Multiple positive prompts detected, splitting by '|'")
positive_prompts_raw = [p.strip() for p in positive_prompt.split('|')]
elif "[1]" in positive_prompt:
log.info("Multiple positive prompts detected, splitting by [#] and enabling EchoShot")
import re
segments = re.split(r'\[\d+\]', positive_prompt)
positive_prompts_raw = [segment.strip() for segment in segments if segment.strip()]
assert len(positive_prompts_raw) > 1 and len(positive_prompts_raw) < 7, 'Input shot num must between 2~6 !'
echoshot = True
else:
positive_prompts_raw = [positive_prompt.strip()]
for p in positive_prompts_raw:
cleaned_prompt, weights = self.parse_prompt_weights(p)
positive_prompts.append(cleaned_prompt)
all_weights.append(weights)
mm.soft_empty_cache()
encoder.model.to(device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=True):
@@ -273,6 +287,7 @@ class WanVideoTextEncode:
prompt_embeds_dict = {
"prompt_embeds": context,
"negative_prompt_embeds": context_null,
"echoshot": echoshot,
}
# Save each part to its own cache file if needed
@@ -1582,12 +1597,11 @@ class WanVideoSampler:
audio_scale = 1.0
if fantasytalking_embeds is not None:
audio_proj = fantasytalking_embeds["audio_proj"].to(device)
audio_context_lens = fantasytalking_embeds["audio_context_lens"]
audio_scale = fantasytalking_embeds["audio_scale"]
audio_cfg_scale = fantasytalking_embeds["audio_cfg_scale"]
if not isinstance(audio_cfg_scale, list):
audio_cfg_scale = [audio_cfg_scale] * (steps +1)
log.info(f"Audio proj shape: {audio_proj.shape}, audio context lens: {audio_context_lens}")
log.info(f"Audio proj shape: {audio_proj.shape}")
elif multitalk_embeds is not None:
# Handle single or multiple speaker embeddings
audio_features_in = multitalk_embeds.get("audio_features", None)
@@ -1735,6 +1749,17 @@ class WanVideoSampler:
feta_args = None
enhance_enabled = False
# EchoShot https://github.com/D2I-ai/EchoShot
echoshot = False
shot_len = None
if text_embeds is not None:
echoshot = text_embeds.get("echoshot", False)
if echoshot:
shot_num = len(text_embeds["prompt_embeds"])
shot_len = [latent_video_length//shot_num] * (shot_num-1)
shot_len.append(latent_video_length-sum(shot_len))
log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
#region transformer settings
#rope
freqs = None
@@ -2083,7 +2108,6 @@ class WanVideoSampler:
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
'fun_camera': control_camera_input if control_camera_latents is not None else None,
'audio_proj': audio_proj if fantasytalking_embeds is not None else None,
'audio_context_lens': audio_context_lens if fantasytalking_embeds is not None else None,
'audio_scale': audio_scale,
"pcd_data": pcd_data_input,
"controlnet": controlnet,
@@ -2092,6 +2116,7 @@ class WanVideoSampler:
"nag_context": text_embeds.get("nag_prompt_embeds", None),
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
"inner_t": [shot_len] if shot_len else None,
}
batch_size = 1
+68 -43
View File
@@ -26,10 +26,11 @@ __all__ = ['WanModel']
from tqdm import tqdm
import gc
import comfy.model_management as mm
from comfy import model_management as mm
from ...utils import log, get_module_memory_mb
from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCacheState, relative_l1_distance
from ...multitalk.multitalk import get_attn_map_with_target
from ...echoshot.echoshot import rope_apply_z, rope_apply_c
from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
@@ -432,10 +433,10 @@ class WanSelfAttention(nn.Module):
k_negative = self.norm_k(self.k(context_negative)).view(b, -1, n, d)
v_negative = self.v(context_negative).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, k_lens=None, attention_mode=self.attention_mode)
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode)
x_positive = x_positive.flatten(2)
x_negative = attention(q, k_negative, v_negative, k_lens=None, attention_mode=self.attention_mode)
x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode)
x_negative = x_negative.flatten(2)
nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1)
@@ -453,15 +454,15 @@ class WanSelfAttention(nn.Module):
return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
#region T2V crossattn
#region crossattn
class WanT2VCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps)
self.attention_mode = attention_mode
def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy"):
def forward(self, x, context, grid_sizes, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy", inner_t=None, inner_c=None, cross_freqs=None):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
@@ -471,7 +472,13 @@ class WanT2VCrossAttention(WanSelfAttention):
else:
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, k_lens=None, attention_mode=self.attention_mode)
#EchoShot rope
if inner_t is not None and cross_freqs is not None and not is_uncond:
q = rope_apply_z(q, grid_sizes, cross_freqs, inner_t).to(q)
k = rope_apply_c(k, cross_freqs, inner_c).to(q)
x_text = attention(q, k, v, attention_mode=self.attention_mode)
x_text = x_text.flatten(2)
x = x_text
@@ -483,13 +490,13 @@ class WanT2VCrossAttention(WanSelfAttention):
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
audio_x = attention(
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode
audio_q, ip_key, ip_value, attention_mode=self.attention_mode
)
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode).flatten(2)
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
x = x + audio_x * audio_scale
@@ -506,13 +513,13 @@ class WanI2VCrossAttention(WanSelfAttention):
self.norm_k_img = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
def forward(self, x, context, context_lens, clip_embed, audio_proj=None, audio_context_lens=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy"):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
**kwargs):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query
@@ -524,13 +531,13 @@ class WanI2VCrossAttention(WanSelfAttention):
# text attention
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode).flatten(2)
x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
#img attention
if clip_embed is not None:
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).flatten(2)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
x = x_text + img_x
else:
x = x_text
@@ -542,14 +549,14 @@ class WanI2VCrossAttention(WanSelfAttention):
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
audio_x = attention(
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode
audio_q, ip_key, ip_value, attention_mode=self.attention_mode
)
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode).flatten(2)
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
x = x + audio_x * audio_scale
x = self.o(x)
@@ -651,14 +658,12 @@ class WanAttentionBlock(nn.Module):
grid_sizes,
freqs,
context,
context_lens,
current_step,
last_step=False,
video_attention_split_steps=[],
clip_embed=None,
camera_embed=None,
audio_proj=None,
audio_context_lens=None,
audio_scale=1.0,
num_latent_frames=21,
enhance_enabled=False,
@@ -668,7 +673,10 @@ class WanAttentionBlock(nn.Module):
is_uncond=False,
multitalk_audio_embedding=None,
ref_target_masks=None,
human_num=0
human_num=0,
inner_t=None,
inner_c=None,
cross_freqs=None,
):
r"""
Args:
@@ -710,7 +718,11 @@ class WanAttentionBlock(nn.Module):
k=rope_apply(k, grid_sizes, freqs)
#self-attention
split_attn = context is not None and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1
split_attn = (context is not None
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
and x.shape[0] == 1
and inner_t is not None
)
if split_attn:
y = self.self_attn.forward_split(
q, k, v,
@@ -755,12 +767,11 @@ class WanAttentionBlock(nn.Module):
if split_attn:
if nag_context is not None:
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
x = self.split_cross_attn_ffn(x, context, context_lens, shift_mlp, scale_mlp, gate_mlp, clip_embed=clip_embed, grid_sizes=grid_sizes)
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
else:
x = self.cross_attn_ffn(x, context, context_lens, shift_mlp, scale_mlp, gate_mlp, clip_embed=clip_embed, grid_sizes=grid_sizes,
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
multitalk_audio_embedding=multitalk_audio_embedding, x_ref_attn_map=x_ref_attn_map, human_num=human_num)
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs)
else:
if self.rope_func == "comfy_chunked":
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
@@ -771,12 +782,14 @@ class WanAttentionBlock(nn.Module):
return x
def cross_attn_ffn(self, x, context, context_lens, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None,
audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, nag_params={},
nag_context=None, is_uncond=False, multitalk_audio_embedding=None, x_ref_attn_map=None, human_num=0):
x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed,
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, rope_func=self.rope_func)
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
audio_proj, audio_scale, num_latent_frames, nag_params,
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs):
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
audio_proj=audio_proj, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs)
#multitalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
@@ -791,7 +804,7 @@ class WanAttentionBlock(nn.Module):
return x
@torch.compiler.disable()
def split_cross_attn_ffn(self, x, context, context_lens, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
# Get number of prompts
num_prompts = context.shape[0]
num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0]
@@ -820,9 +833,6 @@ class WanAttentionBlock(nn.Module):
# 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
@@ -834,7 +844,7 @@ class WanAttentionBlock(nn.Module):
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 = self.cross_attn(self.norm3(x_segment), segment_context, clip_embed=segment_clip_embed)
processed_segment = processed_segment.to(x.dtype)
# Add to combined result
@@ -1354,7 +1364,6 @@ class WanModel(ModelMixin, ConfigMixin):
fun_ref=None,
fun_camera=None,
audio_proj=None,
audio_context_lens=None,
audio_scale=1.0,
pcd_data=None,
controlnet=None,
@@ -1363,7 +1372,8 @@ class WanModel(ModelMixin, ConfigMixin):
nag_params={},
nag_context=None,
multitalk_audio=None,
ref_target_masks=None
ref_target_masks=None,
inner_t=None,
):
r"""
Forward pass through the diffusion model
@@ -1509,9 +1519,16 @@ class WanModel(ModelMixin, ConfigMixin):
else:
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
freqs = self.rope_embedder(img_ids).movedim(1, 2)
self.cached_freqs = freqs
self.cached_shape = current_shape
self.cached_cond = has_cond
# EchoShot cross attn freqs
inner_c = None
if inner_t is not None:
d = self.dim // self.num_heads
self.cross_freqs = rope_params(100, d).to(device=x.device)
# time embeddings
if t.dim() == 2:
@@ -1547,17 +1564,24 @@ class WanModel(ModelMixin, ConfigMixin):
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
#context (text embedding)
context_lens = None
if hasattr(self, "text_embedding") and context != []:
if self.offload_txt_emb:
self.text_embedding.to(self.main_device)
if inner_t is not None:
if nag_context is not None:
raise NotImplementedError("nag_context is not supported with EchoShot")
inner_c = [[u.shape[0] for u in context]]
context = self.text_embedding(
torch.stack([
torch.cat(
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]).to(x.dtype))
# NAG
if nag_context is not None:
nag_context = self.text_embedding(
@@ -1737,14 +1761,12 @@ class WanModel(ModelMixin, ConfigMixin):
grid_sizes=grid_sizes,
freqs=freqs,
context=context,
context_lens=context_lens,
clip_embed=clip_embed,
current_step=current_step,
last_step=last_step,
video_attention_split_steps=self.video_attention_split_steps,
camera_embed=camera_embed,
audio_proj=audio_proj,
audio_context_lens=audio_context_lens,
num_latent_frames = F,
enhance_enabled=enhance_enabled,
audio_scale=audio_scale,
@@ -1754,7 +1776,10 @@ class WanModel(ModelMixin, ConfigMixin):
is_uncond = is_uncond,
multitalk_audio_embedding=multitalk_audio_embedding if multitalk_audio is not None else None,
ref_target_masks=token_ref_target_masks if multitalk_audio is not None else None,
human_num=human_num if multitalk_audio is not None else 0
human_num=human_num if multitalk_audio is not None else 0,
inner_t=inner_t,
inner_c=inner_c,
cross_freqs=self.cross_freqs if inner_t is not None else None,
)
if vace_data is not None: