Don't trigger context windows if not enough frames

This commit is contained in:
kijai
2025-12-15 18:02:41 +02:00
parent 4dbba4d06d
commit e6bd1b413a
+45 -41
View File
@@ -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)