Merge branch 'main' into dev

This commit is contained in:
kijai
2025-08-25 12:51:46 +03:00
8 changed files with 317 additions and 188 deletions
+1 -1
View File
@@ -100,7 +100,7 @@ class WanVideoEasyCache:
return {
"required": {
"easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}),
"start_step": ("INT", {"default": 10, "min": 1, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}),
"start_step": ("INT", {"default": 10, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
+32 -1
View File
@@ -43,7 +43,38 @@ def get_emo_feature(frame_list, face_aligner, pd_fpg_motion, device):
comfy_pbar = ProgressBar(3)
_, landmark_list, rect_list = det_landmarks(face_aligner, frame_list, comfy_pbar)
_, landmark_list, rect_list = det_landmarks(face_aligner, frame_list, comfy_pbar)
# Fill missing landmarks and rects with previous valid one
last_valid_landmark = None
last_valid_rect = None
for i in range(len(landmark_list)):
if landmark_list[i] is None:
landmark_list[i] = last_valid_landmark
else:
last_valid_landmark = landmark_list[i]
if rect_list[i] is None:
rect_list[i] = last_valid_rect
else:
last_valid_rect = rect_list[i]
# Forward fill for leading None values
if landmark_list[0] is None:
first_valid = next((l for l in landmark_list if l is not None), None)
for i in range(len(landmark_list)):
if landmark_list[i] is None:
landmark_list[i] = first_valid
else:
break
if rect_list[0] is None:
first_valid = next((r for r in rect_list if r is not None), None)
for i in range(len(rect_list)):
if rect_list[i] is None:
rect_list[i] = first_valid
else:
break
emo_list = get_drive_expression_pd_fgc(pd_fpg_motion, frame_list, landmark_list, device)
comfy_pbar.update(1)
+23 -17
View File
@@ -156,7 +156,7 @@ def det_landmarks(face_aligner, frame_list, comfy_pbar):
face_aligner.reset_track()
with tqdm(total=len(frame_list)) as pbar:
for frame in frame_list:
for i, frame in enumerate(frame_list):
faces = face_aligner.forward(frame)
if len(faces) > 0:
face = sorted(
@@ -167,36 +167,42 @@ def det_landmarks(face_aligner, frame_list, comfy_pbar):
rect_list.append(face["face_rect"])
new_frame_list.append(frame)
else:
log.warning(f"No face detected in the frame {frame}, skipping.")
log.warning(f"No face detected in the frame {i}, inserting empty frame.")
rect_list.append(None) # Add placeholder
new_frame_list.append(None) # Add placeholder
pbar.set_description("DET stage1")
pbar.update()
comfy_pbar.update(1)
assert len(new_frame_list) > 0
face_aligner.reset_track()
save_frame_list = []
save_landmark_list = []
with tqdm(total=len(new_frame_list)) as pbar:
for frame, rect in zip(new_frame_list, rect_list):
faces = face_aligner.forward(frame, pre_rect=rect)
if len(faces) > 0:
face = sorted(
faces,
key=lambda x: (x["face_rect"][2] - x["face_rect"][0])
* (x["face_rect"][3] - x["face_rect"][1]),
)[-1]
landmarks = face["pre_kpt_222"]
save_frame_list.append(frame)
save_landmark_list.append(landmarks)
for i, (frame, rect) in enumerate(zip(new_frame_list, rect_list)):
if frame is None or rect is None:
save_frame_list.append(None)
save_landmark_list.append(None)
log.warning(f"No face detected in the frame {i}, inserting empty landmark.")
else:
log.warning(f"No face detected in the frame {frame}, skipping.")
faces = face_aligner.forward(frame, pre_rect=rect)
if len(faces) > 0:
face = sorted(
faces,
key=lambda x: (x["face_rect"][2] - x["face_rect"][0])
* (x["face_rect"][3] - x["face_rect"][1]),
)[-1]
landmarks = face["pre_kpt_222"]
save_frame_list.append(frame)
save_landmark_list.append(landmarks)
else:
save_frame_list.append(None)
save_landmark_list.append(None)
log.warning(f"No face detected in the frame {i}, inserting empty landmark.")
pbar.set_description("DET stage2")
pbar.update()
comfy_pbar.update(1)
assert len(save_frame_list) > 0
save_landmark_list = np.stack(save_landmark_list, axis=0)
face_aligner.reset_track()
return save_frame_list, save_landmark_list, rect_list
+2 -2
View File
@@ -78,8 +78,8 @@ class MultiTalkWav2VecEmbeds:
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize the audio loudness to -23 LUFS"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
"multi_audio_type": (["para", "add"], {"default": "para", "tooltip": "'para' overlay speakers in parallel, 'add' concatenate sequentially"}),
},
"optional" : {
+163 -152
View File
@@ -912,40 +912,57 @@ class WanVideoImageToVideoEncode:
# Resize and rearrange the input image dimensions
if start_image is not None:
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
start_image = start_image[..., :3]
if start_image.shape[1] != H or start_image.shape[2] != W:
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
else:
resized_start_image = start_image.permute(3, 0, 1, 2) # C, T, H, W
resized_start_image = resized_start_image * 2 - 1
if noise_aug_strength > 0.0:
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
if end_image is not None:
resized_end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
end_image = end_image[..., :3]
if end_image.shape[1] != H or end_image.shape[2] != W:
resized_end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
else:
resized_end_image = end_image.permute(3, 0, 1, 2) # C, T, H, W
resized_end_image = resized_end_image * 2 - 1
if noise_aug_strength > 0.0:
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
# Concatenate image with zero frames and encode
vae.to(device)
if temporal_mask is None:
if start_image is not None and end_image is None:
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device)
concatenated = torch.cat([resized_start_image.to(device), zero_frames], dim=1)
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
del resized_start_image, zero_frames
elif start_image is None and end_image is not None:
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device)
concatenated = torch.cat([zero_frames, resized_end_image.to(device)], dim=1)
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del zero_frames
elif start_image is None and end_image is None:
concatenated = torch.zeros(3, num_frames, H, W, device=device)
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
else:
if fun_or_fl2v_model:
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device)
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
else:
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device)
concatenated = torch.cat([resized_start_image.to(device), zero_frames, resized_end_image.to(device)], dim=1)
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del resized_start_image, zero_frames
else:
temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1)
concatenated = resized_start_image[:,:num_frames] * temporal_mask[:num_frames].unsqueeze(0)
del resized_start_image, temporal_mask
mm.soft_empty_cache()
gc.collect()
vae.to(device)
y = vae.encode([concatenated], device, end_=(end_image is not None and not fun_or_fl2v_model),tiled=tiled_vae)[0]
vae.model.clear_cache()
del concatenated
y = vae.encode([concatenated.to(device=device, dtype=vae.dtype)], device, end_=(end_image is not None and not fun_or_fl2v_model),tiled=tiled_vae)[0]
has_ref = False
if extra_latents is not None:
samples = extra_latents["samples"].squeeze(0)
@@ -963,8 +980,7 @@ class WanVideoImageToVideoEncode:
if add_cond_latents is not None:
add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
vae.model.clear_cache()
if force_offload:
vae.model.to(offload_device)
mm.soft_empty_cache()
@@ -1160,9 +1176,9 @@ class WanVideoControlEmbeds:
return {"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the control signal"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the control signal"}),
"latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}),
},
"optional": {
"latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}),
"fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}),
}
}
@@ -1172,10 +1188,9 @@ class WanVideoControlEmbeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, start_percent, end_percent, fun_ref_image=None, latents=None):
if latents is not None:
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
def process(self, latents, start_percent, end_percent, fun_ref_image=None):
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
num_frames = (T - 1) * 4 + 1
seq_len = math.ceil((H * W) / 4 * ((num_frames - 1) // 4 + 1))
@@ -1363,18 +1378,20 @@ class WanVideoVACEEncode:
else:
assert len(frames) == len(ref_images)
pbar = ProgressBar(len(frames))
if masks is None:
latents = vae.encode(frames, device=device, tiled=tiled_vae)
else:
inactive = [i * (1 - m) + 0 * m for i, m in zip(frames, masks)]
reactive = [i * m + 0 * (1 - m) for i, m in zip(frames, masks)]
del frames
inactive = vae.encode(inactive, device=device, tiled=tiled_vae)
reactive = vae.encode(reactive, device=device, tiled=tiled_vae)
latents = [torch.cat((u, c), dim=0) for u, c in zip(inactive, reactive)]
del inactive, reactive
vae.model.clear_cache()
cat_latents = []
pbar = ProgressBar(len(frames))
for latent, refs in zip(latents, ref_images):
if refs is not None:
if masks is None:
@@ -1578,7 +1595,7 @@ class WanVideoScheduler: #WIP
return (scheduler,)
rope_functions = ["default", "comfy", "comfy_chunked"]
class WanVideoRoPEFunction: #WIP
class WanVideoRoPEFunction:
@classmethod
def INPUT_TYPES(s):
return {"required": {
@@ -1727,7 +1744,7 @@ class WanVideoSampler:
#region Scheduler
sample_scheduler = None
if scheduler != "multitalk":
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
log.info(f"sigmas: {sample_scheduler.sigmas}")
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
@@ -1746,32 +1763,15 @@ class WanVideoSampler:
noise_pred_flipped = None
if isinstance(cfg, list):
if steps != len(cfg):
log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Setting step count to match.")
steps = len(cfg)
if steps < len(cfg):
log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Slicing cfg list to match steps.")
cfg = cfg[:steps]
elif steps > len(cfg):
log.info(f"Received only {len(cfg)} cfg values, but {steps} steps. Extending cfg list to match steps.")
cfg.extend([cfg[-1]] * (steps - len(cfg)))
log.info(f"Using per-step cfg list: {cfg}")
else:
cfg = [cfg] * (steps + 1)
if end_step != -1:
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
log.info(f"Sampling until step {end_step}, timestep: {timesteps[-1]}")
if start_step > 0:
timesteps = timesteps[start_step:]
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
log.info(f"Skipping first {start_step} steps, starting from timestep {timesteps[0]}")
log.info(f"timesteps: {timesteps}")
if sample_scheduler is not None:
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
scheduler_step_args = {"generator": seed_g}
step_sig = inspect.signature(sample_scheduler.step)
for arg in list(scheduler_step_args.keys()):
if arg not in step_sig.parameters:
scheduler_step_args.pop(arg)
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = None
vace_data = vace_context = vace_scale = None
@@ -2282,19 +2282,32 @@ class WanVideoSampler:
transformer.to(device)
# Initialize Cache if enabled
previous_cache_states = None
transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False
cache_args = teacache_args if teacache_args is not None else cache_args #for backward compatibility on old workflows
if cache_args is not None:
from .cache_methods.cache_methods import set_transformer_cache_method
transformer = set_transformer_cache_method(transformer, timesteps, cache_args)
from .cache_methods.cache_methods import set_transformer_cache_method
transformer = set_transformer_cache_method(transformer, timesteps, cache_args)
# Initialize cache state
self.cache_state = [None, None]
if phantom_latents is not None:
log.info(f"Phantom latents shape: {phantom_latents.shape}")
self.cache_state = [None, None, None]
self.cache_state_source = [None, None]
self.cache_states_context = []
# Initialize cache state
if samples is not None:
previous_cache_states = samples.get("cache_states", None)
print("Using previous cache states", previous_cache_states)
if previous_cache_states is not None:
log.info("Using cache states from previous sampler")
self.cache_state = previous_cache_states["cache_state"]
transformer.easycache_state = previous_cache_states["easycache_state"]
transformer.magcache_state = previous_cache_states["magcache_state"]
transformer.teacache_state = previous_cache_states["teacache_state"]
if previous_cache_states is None:
self.cache_state = [None, None]
if phantom_latents is not None:
log.info(f"Phantom latents shape: {phantom_latents.shape}")
self.cache_state = [None, None, None]
self.cache_state_source = [None, None]
self.cache_states_context = []
# Skip layer guidance (SLG)
if slg_args is not None:
@@ -2776,7 +2789,7 @@ class WanVideoSampler:
# FreeInit noise reinitialization (after first iteration)
if freeinit_args is not None and iter_idx > 0:
# restart scheduler for each iteration
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
# Re-apply start_step and end_step logic to timesteps and sigmas
if end_step != -1:
@@ -3166,9 +3179,10 @@ class WanVideoSampler:
audio_end_idx = audio_start_idx + clip_length
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
if multitalk_embeds is not None:
total_frames = len(multitalk_audio_embedding[0])
audio_embedding = multitalk_audio_embedding
human_num = len(audio_embedding)
audio_embs = None
pcd_data = pcd_data_input = None
if uni3c_embeds is not None:
@@ -3182,13 +3196,10 @@ class WanVideoSampler:
"end": uni3c_embeds["end"],
}
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
audio_embedding = multitalk_audio_embedding
human_num = len(audio_embedding)
audio_embs = None
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
@@ -3205,29 +3216,13 @@ class WanVideoSampler:
audio_embs.append(audio_emb)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
required_frames = audio_end_idx - audio_start_idx
if audio_end_idx > num_frames:
pad_len = audio_end_idx - num_frames
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
pcd_data['render_latent'] = render_latent
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
seq_len = ((frame_num - 1) // VAE_STRIDE[0] + 1) * lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(
16, (frame_num - 1) // 4 + 1,
16, latent_frame_num,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
@@ -3280,51 +3275,27 @@ class WanVideoSampler:
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
window_vace_data = None
if vace_data is not None:
window_vace_data = []
for vace_entry in vace_data:
partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
if has_ref:
partial_context[:, 0] = vace_entry["context"][0][:, 0]
window_vace_data.append({
"context": [partial_context],
"scale": vace_entry["scale"],
"start": vace_entry["start"],
"end": vace_entry["end"],
"seq_len": vace_entry["seq_len"]
})
# get image cond mask
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
if mode == "multitalk":
msk[:, cur_motion_frames_num:] = 0
else:
msk[:, 1:] = 0
msk = torch.concat([
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
msk = msk.transpose(1, 2).to(dtype) # B 4 T H W
mm.soft_empty_cache()
# zero padding and vae encode
# zero padding and vae encode for img cond
if cond_image is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "multitalk":
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.model.clear_cache()
vae.to(offload_device)
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
motion_frame_index = cur_motion_frames_num if mode == "multitalk" else 1
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :motion_frame_index] = 1
y = torch.cat([msk, y]) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
@@ -3336,33 +3307,8 @@ class WanVideoSampler:
timesteps = [torch.tensor([t], device=device) for t in timesteps]
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
else:
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
steps = len(timesteps)
if end_step != -1 and start_step >= end_step:
raise ValueError("start_step must be less than end_step")
if denoise_strength < 1.0:
if start_step != 0:
raise ValueError("start_step must be 0 when denoise_strength is used")
start_step = steps - int(steps * denoise_strength) - 1
if end_step != -1:
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
if start_step > 0:
timesteps = timesteps[start_step:]
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
if sample_scheduler is not None:
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
transformed_timesteps = []
for t in timesteps:
t_tensor = torch.tensor([t.item()], device=device)
transformed_timesteps.append(t_tensor)
transformed_timesteps.append(torch.tensor([0.], device=device))
timesteps = transformed_timesteps
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
# sample videos
latent = noise
@@ -3415,6 +3361,43 @@ class WanVideoSampler:
else:
positive = text_embeds["prompt_embeds"]
window_vace_data = None
# if vace_data is not None:
# window_vace_data = []
# for vace_entry in vace_data:
# partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
# if has_ref:
# partial_context[:, 0] = vace_entry["context"][0][:, 0]
# window_vace_data.append({
# "context": [partial_context],
# "scale": vace_entry["scale"],
# "start": vace_entry["start"],
# "end": vace_entry["end"],
# "seq_len": vace_entry["seq_len"]
# })
# uni3c slices
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
required_frames = audio_end_idx - audio_start_idx
if audio_end_idx > num_frames:
pad_len = audio_end_idx - num_frames
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
vae.model.clear_cache()
vae.to(offload_device)
pcd_data['render_latent'] = render_latent
# unianimate slices
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
@@ -3427,6 +3410,22 @@ class WanVideoSampler:
"end_percent": unianimate_poses["end_percent"]
}
# fantasy portrait slices
partial_fantasy_portrait_input = None
if fantasy_portrait_input is not None:
adapter_proj = fantasy_portrait_input["adapter_proj"]
if latent_end_idx > adapter_proj.shape[1]:
pad_len = latent_end_idx - adapter_proj.shape[1]
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
else:
padded_proj = adapter_proj
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
mm.soft_empty_cache()
gc.collect()
# sampling loop
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
@@ -3437,7 +3436,7 @@ class WanVideoSampler:
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
@@ -3472,11 +3471,12 @@ class WanVideoSampler:
else:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, y, msk, latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer)
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.model.clear_cache()
vae.to(offload_device)
sampling_pbar.close()
@@ -3517,7 +3517,6 @@ class WanVideoSampler:
cond_image = cond_
del videos, latent
mm.soft_empty_cache()
# Repeat audio emb
if multitalk_embeds is not None:
@@ -3663,9 +3662,20 @@ class WanVideoSampler:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
raise e
if phantom_latents is not None:
latent = latent[:,:-phantom_latents.shape[1]]
cache_states = None
if cache_args is not None:
cache_report(transformer, cache_args)
if end_step != -1 and end_step < total_steps:
cache_states = {
"cache_state": self.cache_state,
"easycache_state": transformer.easycache_state,
"teacache_state": transformer.teacache_state,
"magcache_state": transformer.magcache_state,
}
if force_offload:
if not model["auto_cpu_offload"]:
@@ -3685,7 +3695,8 @@ class WanVideoSampler:
"has_ref": has_ref,
"drop_last": drop_last,
"generator_state": seed_g.get_state(),
"original_image": original_image.cpu() if original_image is not None else None
"original_image": original_image.cpu() if original_image is not None else None,
"cache_states": cache_states
},{
"samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None,
})
@@ -3757,7 +3768,7 @@ class WanVideoDecode:
images = images.permute(1, 2, 3, 0).cpu().float()
return (images,)
else:
if end_image is not None:
if end_image:
enable_vae_tiling = False
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
vae.model.clear_cache()
@@ -3878,7 +3889,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
"WanVideoRoPEFunction": WanVideoRoPEFunction
"WanVideoRoPEFunction": WanVideoRoPEFunction,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -3912,5 +3923,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent",
"WanVideoAddControlEmbeds": "WanVideo Add Control Embeds",
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion",
"WanVideoRoPEFunction": "WanVideo RoPE Function"
"WanVideoRoPEFunction": "WanVideo RoPE Function",
}
+23 -4
View File
@@ -391,6 +391,23 @@ class WanVideoLatentReScale:
return (samples,)
class WanVideoSigmaToStep:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"sigma": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.001}),
},
}
RETURN_TYPES = ("INT", )
RETURN_NAMES = ("step",)
FUNCTION = "convert"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Simply passes a float value as an integer, used to set start/end steps with sigma threshold"
def convert(self, sigma):
return (sigma,)
NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
@@ -398,8 +415,9 @@ NODE_CLASS_MAPPINGS = {
"CreateCFGScheduleFloatList": CreateCFGScheduleFloatList,
"DummyComfyWanModelObject": DummyComfyWanModelObject,
"WanVideoLatentReScale": WanVideoLatentReScale,
"CreateScheduleFloatList": CreateScheduleFloatList
}
"CreateScheduleFloatList": CreateScheduleFloatList,
"WanVideoSigmaToStep": WanVideoSigmaToStep
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
"WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame",
@@ -407,5 +425,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"CreateCFGScheduleFloatList": "Create CFG Schedule Float List",
"DummyComfyWanModelObject": "Dummy Comfy Wan Model Object",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
"CreateScheduleFloatList": "Create Schedule Float List"
}
"CreateScheduleFloatList": "Create Schedule Float List",
"WanVideoSigmaToStep": "WanVideo Sigma To Step"
}
+24 -8
View File
@@ -156,7 +156,7 @@ def sinusoidal_embedding_1d(dim, position):
# preprocess
assert dim % 2 == 0
half = dim // 2
position = position.type(torch.float64)
position = position.type(torch.float32)
# calculation
sinusoid = torch.outer(
@@ -2027,28 +2027,34 @@ class WanModel(torch.nn.Module):
previous_raw_output = state.get('previous_raw_output')
cache = state.get('cache')
accumulated_error = state.get('accumulated_error')
k = state.get('k', 1)
if previous_raw_input is not None and previous_raw_output is not None:
raw_input = x.clone()
# Calculate input change
raw_input_change = (raw_input - previous_raw_input.to(raw_input.device)).abs().mean()
accumulated_error += raw_input_change
output_norm = (previous_raw_output.to(x.device)).abs().mean()
combined_pred_change = (raw_input_change / output_norm) * k
accumulated_error += combined_pred_change
# Predict output change
if accumulated_error < self.easycache_thresh:
should_calc = False
x = raw_input + cache.to(x.device)
self.easycache_state.get(pred_id)['skipped_steps'].append(current_step)
state['skipped_steps'].append(current_step)
else:
should_calc = True
accumulated_error = 0.0
else:
should_calc = True
if self.enable_easycache:
original_x = x.clone().to(self.cache_device)
if should_calc:
if self.enable_teacache or self.enable_magcache or self.enable_easycache:
original_x = x.to(self.cache_device).clone()
if self.enable_teacache or self.enable_magcache:
original_x = x.clone().to(self.cache_device)
if hasattr(self, "dwpose_embedding") and unianim_data is not None:
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
@@ -2188,13 +2194,23 @@ class WanModel(torch.nn.Module):
residual_cache=(x.to(original_x.device) - original_x)
)
elif self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
x_out = x.clone().to(original_x.device)
output_change = (x_out - original_x).abs().mean()
input_change = (original_x - x_out).abs().mean()
self.easycache_state.update(
pred_id,
previous_raw_input=original_x,
previous_raw_output=x.clone(),
previous_raw_output=x_out,
cache=x.to(original_x.device) - original_x,
accumulated_error=0.0
k = output_change / input_change,
accumulated_error = 0.0
)
if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
self.easycache_state.update(
pred_id,
previous_raw_output=x.clone(),
)
if self.ref_conv is not None and fun_ref is not None:
fun_ref_length = fun_ref.size(1)
+49 -3
View File
@@ -6,7 +6,7 @@ from .flowmatch_pusa import FlowMatchSchedulerPusa
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler
import numpy as np
import inspect
from ...utils import log
scheduler_list = [
@@ -23,7 +23,7 @@ scheduler_list = [
"multitalk"
]
def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None):
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None, seed_g=None):
timesteps = None
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
@@ -99,4 +99,50 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
if timesteps is None:
timesteps = sample_scheduler.timesteps
return sample_scheduler, timesteps
steps = len(timesteps)
if end_step != -1 and start_step >= end_step:
raise ValueError("start_step must be less than end_step")
if denoise_strength < 1.0:
if start_step != 0:
raise ValueError("start_step must be 0 when denoise_strength is used")
start_step = steps - int(steps * denoise_strength) - 1
# Determine start and end indices for slicing
start_idx = 0
end_idx = len(timesteps) - 1
if isinstance(start_step, float):
idxs = (sample_scheduler.sigmas <= start_step).nonzero(as_tuple=True)[0]
if len(idxs) > 0:
start_idx = idxs[0].item()
elif isinstance(start_step, int):
if start_step > 0:
start_idx = start_step
if isinstance(end_step, float):
idxs = (sample_scheduler.sigmas >= end_step).nonzero(as_tuple=True)[0]
if len(idxs) > 0:
end_idx = idxs[-1].item()
elif isinstance(end_step, int):
if end_step != -1:
end_idx = end_step - 1
# Slice timesteps and sigmas once, based on indices
timesteps = timesteps[start_idx:end_idx+1]
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
log.info(f"timesteps: {timesteps}")
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
if seed_g is not None:
scheduler_step_args = {"generator": seed_g}
step_sig = inspect.signature(sample_scheduler.step)
for arg in list(scheduler_step_args.keys()):
if arg not in step_sig.parameters:
scheduler_step_args.pop(arg)
return sample_scheduler, timesteps, scheduler_step_args