Don't trigger context windows if not enough frames
This commit is contained in:
+45
-41
@@ -680,50 +680,54 @@ class WanVideoSampler:
|
||||
is_looped = False
|
||||
context_reference_latent = None
|
||||
if context_options is not None:
|
||||
context_schedule = context_options["context_schedule"]
|
||||
context_frames = (context_options["context_frames"] - 1) // 4 + 1
|
||||
context_stride = context_options["context_stride"] // 4
|
||||
context_overlap = context_options["context_overlap"] // 4
|
||||
context_reference_latent = context_options.get("reference_latent", None)
|
||||
if context_options["context_frames"] <= num_frames:
|
||||
context_schedule = context_options["context_schedule"]
|
||||
context_frames = (context_options["context_frames"] - 1) // 4 + 1
|
||||
context_stride = context_options["context_stride"] // 4
|
||||
context_overlap = context_options["context_overlap"] // 4
|
||||
context_reference_latent = context_options.get("reference_latent", None)
|
||||
|
||||
# Get total number of prompts
|
||||
num_prompts = len(text_embeds["prompt_embeds"])
|
||||
log.info(f"Number of prompts: {num_prompts}")
|
||||
# Calculate which section this context window belongs to
|
||||
section_size = (latent_video_length / num_prompts) if num_prompts != 0 else 1
|
||||
log.info(f"Section size: {section_size}")
|
||||
is_looped = context_schedule == "uniform_looped"
|
||||
# Get total number of prompts
|
||||
num_prompts = len(text_embeds["prompt_embeds"])
|
||||
log.info(f"Number of prompts: {num_prompts}")
|
||||
# Calculate which section this context window belongs to
|
||||
section_size = (latent_video_length / num_prompts) if num_prompts != 0 else 1
|
||||
log.info(f"Section size: {section_size}")
|
||||
is_looped = context_schedule == "uniform_looped"
|
||||
|
||||
if mocha_embeds is not None:
|
||||
seq_len = (context_frames * 2 + 1 + mocha_num_refs) * (noise.shape[2] * noise.shape[3] // 4)
|
||||
if mocha_embeds is not None:
|
||||
seq_len = (context_frames * 2 + 1 + mocha_num_refs) * (noise.shape[2] * noise.shape[3] // 4)
|
||||
else:
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * context_frames)
|
||||
log.info(f"context window seq len: {seq_len}")
|
||||
|
||||
if context_options["freenoise"]:
|
||||
log.info("Applying FreeNoise")
|
||||
# code from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
|
||||
delta = context_frames - context_overlap
|
||||
for start_idx in range(0, latent_video_length-context_frames, delta):
|
||||
place_idx = start_idx + context_frames
|
||||
if place_idx >= latent_video_length:
|
||||
break
|
||||
end_idx = place_idx - 1
|
||||
|
||||
if end_idx + delta >= latent_video_length:
|
||||
final_delta = latent_video_length - place_idx
|
||||
list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
|
||||
list_idx = list_idx[torch.randperm(final_delta, generator=seed_g)]
|
||||
noise[:, place_idx:place_idx + final_delta, :, :] = noise[:, list_idx, :, :]
|
||||
break
|
||||
list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
|
||||
list_idx = list_idx[torch.randperm(delta, generator=seed_g)]
|
||||
noise[:, place_idx:place_idx + delta, :, :] = noise[:, list_idx, :, :]
|
||||
|
||||
log.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
|
||||
from .context_windows.context import get_context_scheduler, create_window_mask, WindowTracker
|
||||
self.window_tracker = WindowTracker(verbose=context_options["verbose"])
|
||||
context = get_context_scheduler(context_schedule)
|
||||
else:
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * context_frames)
|
||||
log.info(f"context window seq len: {seq_len}")
|
||||
|
||||
if context_options["freenoise"]:
|
||||
log.info("Applying FreeNoise")
|
||||
# code from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
|
||||
delta = context_frames - context_overlap
|
||||
for start_idx in range(0, latent_video_length-context_frames, delta):
|
||||
place_idx = start_idx + context_frames
|
||||
if place_idx >= latent_video_length:
|
||||
break
|
||||
end_idx = place_idx - 1
|
||||
|
||||
if end_idx + delta >= latent_video_length:
|
||||
final_delta = latent_video_length - place_idx
|
||||
list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
|
||||
list_idx = list_idx[torch.randperm(final_delta, generator=seed_g)]
|
||||
noise[:, place_idx:place_idx + final_delta, :, :] = noise[:, list_idx, :, :]
|
||||
break
|
||||
list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
|
||||
list_idx = list_idx[torch.randperm(delta, generator=seed_g)]
|
||||
noise[:, place_idx:place_idx + delta, :, :] = noise[:, list_idx, :, :]
|
||||
|
||||
log.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
|
||||
from .context_windows.context import get_context_scheduler, create_window_mask, WindowTracker
|
||||
self.window_tracker = WindowTracker(verbose=context_options["verbose"])
|
||||
context = get_context_scheduler(context_schedule)
|
||||
log.info("Context frames is larger than total num_frames, disabling context windows")
|
||||
context_options = None
|
||||
|
||||
#MTV Crafter
|
||||
mtv_input = image_embeds.get("mtv_crafter_motion", None)
|
||||
|
||||
Reference in New Issue
Block a user