Merge branch 'main' into dev
This commit is contained in:
@@ -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"}),
|
||||
},
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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" : {
|
||||
|
||||
@@ -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
@@ -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"
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user