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:
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user