Make start/end step work with vid2vid

This commit is contained in:
kijai
2025-07-30 15:59:22 +03:00
parent 406643f8ee
commit ec066008a8
2 changed files with 27 additions and 23 deletions
+23 -17
View File
@@ -1336,8 +1336,6 @@ class WanVideoSampler:
multitalk_sampling = image_embeds.get("multitalk_sampling", False)
if not multitalk_sampling and scheduler == "multitalk":
raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node")
steps = int(steps/denoise_strength)
if text_embeds == None:
text_embeds = {
@@ -1347,30 +1345,42 @@ class WanVideoSampler:
else:
text_embeds = dict_to_device(text_embeds, device)
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)
else:
cfg = [cfg] * (steps +1)
seed_g = torch.Generator(device=torch.device("cpu"))
seed_g.manual_seed(seed)
# Scheduler
#region Scheduler
if scheduler != "multitalk":
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
log.info(f"sigmas: {sample_scheduler.sigmas}")
steps = len(timesteps)
if end_step != -1 or end_step >= steps:
if denoise_strength < 1.0:
start_step = steps - int(steps * denoise_strength) - 1
first_sampler = (end_step != -1 or end_step >= steps)
add_noise_to_samples = True if first_sampler else False
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)
else:
cfg = [cfg] * (steps + 1)
if first_sampler:
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
elif start_step > 0:
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 hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
@@ -1379,10 +1389,6 @@ class WanVideoSampler:
for arg in list(scheduler_step_args.keys()):
if arg not in step_sig.parameters:
scheduler_step_args.pop(arg)
if denoise_strength < 1.0:
steps = int(steps * denoise_strength)
timesteps = timesteps[-(steps + 1):]
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
@@ -1698,7 +1704,7 @@ class WanVideoSampler:
if input_samples.shape[1] != noise.shape[1]:
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
if denoise_strength < 1.0:
if add_noise_to_samples:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
+4 -6
View File
@@ -661,14 +661,12 @@ class WanVideoSetLoRAs:
patcher = model.clone()
lora_low_mem_load = merge_loras = False
merge_loras = False
for l in lora:
lora_low_mem_load = l.get("low_mem_load", False)
merge_loras = l.get("merge_loras", True)
if lora_low_mem_load is True or merge_loras is True:
raise ValueError("Set LoRA node does not use low_mem_load and can't merge LoRAs, disable low_mem_load when and merge_loras in the LoRA select node.")
if merge_loras is True:
raise ValueError("Set LoRA node does not use low_mem_load and can't merge LoRAs, disable 'merge_loras' in the LoRA select node.")
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
@@ -1241,7 +1239,7 @@ class WanVideoModelLoader:
for model in mm.current_loaded_models:
if model._model() == patcher:
mm.current_loaded_models.remove(model)
mm.current_loaded_models.remove(model)
return (patcher,)