context windows
This commit is contained in:
+15
-7
@@ -1172,13 +1172,14 @@ class WanVideoSampler:
|
||||
latent = add_noise(ttm_reference_latents, noise, timesteps[ttm_start_step].to(noise.device)).to(latent)
|
||||
|
||||
# SteadyDancer
|
||||
sdance_embeds = image_embeds.get("sdance_embeds", None)
|
||||
sdancer_input = None
|
||||
if sdance_embeds is not None:
|
||||
print("Using SteadyDancer embeddings")
|
||||
print(f"SteadyDancer embeds keys: {list(sdance_embeds.keys())}")
|
||||
sdancer_input = sdance_embeds.copy()
|
||||
sdancer_input = dict_to_device(sdancer_input, device, dtype)
|
||||
sdancer_embeds = image_embeds.get("sdancer_embeds", None)
|
||||
sdancer_data = sdancer_input = None
|
||||
if sdancer_embeds is not None:
|
||||
log.info("Using SteadyDancer embeddings:")
|
||||
for k, v in sdancer_embeds.items():
|
||||
log.info(f" {k}: {v.shape if isinstance(v, torch.Tensor) else v}")
|
||||
sdancer_data = sdancer_embeds.copy()
|
||||
sdancer_data = dict_to_device(sdancer_data, device, dtype)
|
||||
|
||||
#region model pred
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
||||
@@ -1357,6 +1358,13 @@ class WanVideoSampler:
|
||||
else:
|
||||
uni3c_data_input = uni3c_data
|
||||
|
||||
if context_window is not None and sdancer_data is not None and sdancer_data["cond_pos"].shape[1] != context_frames:
|
||||
sdancer_input = sdancer_data.copy()
|
||||
sdancer_input["cond_pos"] = sdancer_data["cond_pos"][:, context_window]
|
||||
sdancer_input["cond_neg"] = sdancer_data["cond_neg"][:, context_window] if sdancer_data.get("cond_neg", None) is not None else None
|
||||
else:
|
||||
sdancer_input = sdancer_data
|
||||
|
||||
if s2v_pose is not None:
|
||||
if not ((s2v_pose_start_percent <= current_step_percentage <= s2v_pose_end_percent) or \
|
||||
(s2v_pose_end_percent > 0 and idx == 0 and current_step_percentage >= s2v_pose_start_percent)):
|
||||
|
||||
@@ -39,7 +39,7 @@ class WanVideoAddSteadyDancerEmbeds:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, pose_latents_positive, pose_strength_spatial, pose_strength_temporal, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None):
|
||||
sdance_embeds = {
|
||||
sdancer_embeds = {
|
||||
"cond_pos": pose_latents_positive["samples"][0],
|
||||
"cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None,
|
||||
"pose_strength_spatial": pose_strength_spatial,
|
||||
@@ -50,7 +50,7 @@ class WanVideoAddSteadyDancerEmbeds:
|
||||
}
|
||||
|
||||
updated = dict(embeds)
|
||||
updated["sdance_embeds"] = sdance_embeds
|
||||
updated["sdancer_embeds"] = sdancer_embeds
|
||||
return (updated,)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user