MoCha: experimental context windows support

This commit is contained in:
kijai
2025-10-21 18:32:40 +03:00
parent 56120d633e
commit 4e36aee658
2 changed files with 19 additions and 5 deletions
+1 -1
View File
@@ -167,7 +167,7 @@ class MochaEmbeds:
"mocha_embeds": mocha_embeds,
"num_frames": F,
"target_shape": target_shape,
"num_refs": num_refs,
"mocha_num_refs": num_refs,
}
return (image_embeds,)
+18 -4
View File
@@ -418,6 +418,7 @@ class WanVideoSampler:
if image_embeds.get("mocha_embeds", None) is not None:
mocha_embeds = image_embeds.get("mocha_embeds", None)
mocha_num_refs = image_embeds.get("mocha_num_refs", 0)
orig_noise_len = noise.shape[1]
seq_len = image_embeds.get("seq_len", seq_len)
log.info(f"MoCha embeds shape: {mocha_embeds.shape}")
@@ -683,7 +684,11 @@ class WanVideoSampler:
log.info(f"Section size: {section_size}")
is_looped = context_schedule == "uniform_looped"
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * context_frames)
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")
@@ -1208,7 +1213,12 @@ class WanVideoSampler:
z = torch.cat([z, recam_latents.to(z)], dim=1)
if mocha_embeds is not None:
z = torch.cat([z, mocha_embeds.to(z)], dim=1)
if context_window is not None and mocha_embeds.shape[2] != context_frames:
partial_mocha_embeds = mocha_embeds[:, context_window]
partial_mocha_embeds[:, -mocha_num_refs] = mocha_embeds[:, -mocha_num_refs]
z = torch.cat([z, partial_mocha_embeds.to(z)], dim=1)
else:
z = torch.cat([z, mocha_embeds.to(z)], dim=1)
if mtv_input is not None:
if ((mtv_start_percent <= current_step_percentage <= mtv_end_percent) or \
@@ -1991,7 +2001,8 @@ class WanVideoSampler:
partial_timestep[:, :1] = 0
else:
partial_timestep = timestep
#print("Partial timestep:", partial_timestep)
orig_model_input_frames = partial_latent_model_input.shape[1]
noise_pred_context, _, new_teacache = predict_with_cfg(
partial_latent_model_input,
@@ -2007,7 +2018,10 @@ class WanVideoSampler:
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped, window_type=context_options["fuse_method"])
if mocha_embeds is not None:
noise_pred_context = noise_pred_context[:, :orig_model_input_frames]
window_mask = create_window_mask(noise_pred_context, c, noise.shape[1], context_overlap, looped=is_looped, window_type=context_options["fuse_method"])
noise_pred[:, c] += noise_pred_context * window_mask
counter[:, c] += window_mask
context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, len(timesteps))