Closer to original EchoShot implementation, bugfixes

Turns out I never had the right way of using EchoShot and was always just accidentally using my existing rope splitting... which just worked with the EchoShot weights, this commit adds proper EchoShot and it's used when their prompt format is used, the previous behaviour can be restored by using my splitting method with " | " between the prompts. EchoShot example has been updated to reflect that.
This commit is contained in:
kijai
2025-08-07 12:08:42 +03:00
parent 6191451e2c
commit ab06ac2f64
4 changed files with 940 additions and 411 deletions
+39
View File
@@ -60,6 +60,45 @@ def rope_apply_c(x, freqs, inner_c, shift=6):
# 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()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_echoshot(x, grid_sizes, freqs, inner_t, shift=4):
n, c = x.size(2), x.size(3) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# 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 += list(range(shot_ind * shift + s, shot_ind * shift + e))
t_freqs = freqs[0][freq_select]
freqs_i = torch.cat([
# freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
t_freqs.view(f, 1, 1, -1).expand(f, h, w, -1), ###
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-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()
File diff suppressed because it is too large Load Diff
+4 -2
View File
@@ -1965,7 +1965,8 @@ class WanVideoSampler:
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"EchoShot - Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
rope_function = "default" #echoshot does not support comfy rope function
log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
#region transformer settings
#rope
@@ -2991,7 +2992,7 @@ class WanVideoSampler:
# cache generated samples
videos = torch.stack(videos).cpu() # B C T H W
if colormatch != "disabled":
videos = videos[0].permute(1, 2, 3, 0).cpu().numpy()
videos = videos[0].permute(1, 2, 3, 0).cpu().float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
@@ -3236,6 +3237,7 @@ class WanVideoDecode:
if is_looped:
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
temp_images = temp_images.cpu().float()
temp_images = (temp_images - temp_images.min()) / (temp_images.max() - temp_images.min())
images = torch.cat([temp_images[:, 9:].to(images), images[:, 5:]], dim=1)
+24 -24
View File
@@ -27,8 +27,9 @@ 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 ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from comfy.model_management import get_torch_device, get_autocast_device
from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4):
@@ -155,7 +156,6 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs
from comfy.model_management import get_torch_device, get_autocast_device
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply(x, grid_sizes, freqs):
@@ -359,33 +359,30 @@ class WanSelfAttention(nn.Module):
# Split by frames if multiple prompts are provided
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()
# Extract frame, height, width from grid_sizes
frames = grid_sizes[0][0]
height = grid_sizes[0][1]
width = grid_sizes[0][2]
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
actual_chunks = torch.min(torch.tensor(seq_chunks, device=frames.device), frames)
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
# Calculate all chunk boundaries
chunk_indices = torch.arange(actual_chunks, device=frames.device)
chunk_sizes = base_frames_per_chunk + (chunk_indices < extra_frames).long()
chunk_starts = torch.cumsum(torch.cat([torch.zeros(1, device=frames.device), chunk_sizes[:-1]]), dim=0).long()
chunk_ends = chunk_starts + chunk_sizes
# 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)
# Process each chunk using tensor indexing
for i in range(actual_chunks.item()):
start_frame = chunk_starts[i]
end_frame = chunk_ends[i]
# Convert to token indices using tensor operations
start_idx = start_frame * tokens_per_frame
end_idx = end_frame * tokens_per_frame
chunk_q = q[:, start_idx:end_idx, :, :]
chunk_k = k[:, start_idx:end_idx, :, :]
@@ -706,7 +703,10 @@ class WanAttentionBlock(nn.Module):
feta_scores = get_feta_scores(q, k)
#RoPE
if self.rope_func == "comfy":
if inner_t is not None:
q=rope_apply_echoshot(q, grid_sizes, freqs, inner_t).to(q)
k=rope_apply_echoshot(k, grid_sizes, freqs, inner_t).to(k)
elif self.rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
elif self.rope_func == "comfy_chunked":
q, k = apply_rope_comfy_chunked(q, k, freqs)