Update wan_video_v2v_combine.py

This commit is contained in:
smthemex
2025-07-21 21:52:39 +08:00
committed by GitHub
parent 95a351a81b
commit 1d6b89728c
@@ -4,6 +4,7 @@ from ..models.wan_video_text_encoder import WanTextEncoder
from ..models.wan_video_vae import WanVideoVAE
from ..models.wan_video_image_encoder import WanImageEncoder
from ..schedulers.flow_match import FlowMatchScheduler,WanStepDistillScheduler
from ..schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
from .base import BasePipeline
from ..prompters import WanPrompter_ as WanPrompter
import torch, os
@@ -141,14 +142,15 @@ tokenizer_path_=os.path.join(folder_paths.base_path,"custom_nodes\ComfyUI_AniCra
class WanMovieCrafterCombineVideoPipeline_(BasePipeline):
def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=tokenizer_path_, config=None):
def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=tokenizer_path_):
super().__init__(device=device, torch_dtype=torch_dtype)
self.scheduler = WanStepDistillScheduler(config) if config is not None else FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
self.uselcm = False if config is None else True
#self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
# self.scheduler = WanStepDistillScheduler(config) if config is not None else FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
self.uselcm = False
self.useunipc=False
#self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
self.scheduler =FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
self.prompter : WanPrompter= None
self.text_encoder: WanTextEncoder = None
self.image_encoder: WanImageEncoder = None
@@ -320,14 +322,13 @@ class WanMovieCrafterCombineVideoPipeline_(BasePipeline):
@staticmethod
def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None, config=None):
def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None, ):
if device is None: device = model_manager.device
if torch_dtype is None: torch_dtype = model_manager.torch_dtype
pipe = WanMovieCrafterCombineVideoPipeline_(
device=device,
torch_dtype=torch_dtype,
config=config
)
pipe.fetch_models(model_manager)
#model_manager._clean_lora_cache()
@@ -484,8 +485,10 @@ class WanMovieCrafterCombineVideoPipeline_(BasePipeline):
tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
# Scheduler
if self.uselcm :
#self.scheduler.denoising_step_list=get_denoising_step_list(num_inference_steps) #动态调整list
if self.useunipc:
self.scheduler.set_timesteps(num_inference_steps, device=self.device)
elif self.uselcm :
#self.scheduler.denoising_step_list=get_denoising_step_list(num_inference_steps) #不需要动态调整list
self.scheduler.set_denoising_timesteps(device=self.device)
else:
self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift)
@@ -582,23 +585,21 @@ class WanMovieCrafterCombineVideoPipeline_(BasePipeline):
# condition
condition = blend_data + smplx_data
condition = rearrange(condition, 'b c f h w -> b (f h w) c').contiguous()
if not self.uselcm:
if self.useunipc:
for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
# Inference
model_input = latents
noise_pred_posi = model_fn_wan_video(self.dit, model_input, timestep=timestep, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi, add_condition = condition)
## noise_pred_posi = model_fn_wan_video(self.dit, latents, timestep=timestep, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi)
noise_pred_posi = model_fn_wan_video(self.dit, model_input, timestep=timestep, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi, add_condition=condition)
if cfg_scale != 1.0:
noise_pred_nega = model_fn_wan_video(self.dit, latents, timestep=timestep, **prompt_emb_nega, **image_emb, **extra_input, **tea_cache_nega)
noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
else:
noise_pred = noise_pred_posi
# Scheduler
latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
else:
# UniPC step
latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents).prev_sample
elif self.uselcm:
for progress_id in progress_bar_cmd(range(len(self.scheduler.sigmas))):
timestep = self.scheduler.timesteps[progress_id]
# Inference
@@ -617,6 +618,24 @@ class WanMovieCrafterCombineVideoPipeline_(BasePipeline):
self.scheduler.step_post()
latents = self.scheduler.latents
else:
for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
# Inference
model_input = latents
noise_pred_posi = model_fn_wan_video(self.dit, model_input, timestep=timestep, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi, add_condition = condition)
## noise_pred_posi = model_fn_wan_video(self.dit, latents, timestep=timestep, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi)
if cfg_scale != 1.0:
noise_pred_nega = model_fn_wan_video(self.dit, latents, timestep=timestep, **prompt_emb_nega, **image_emb, **extra_input, **tea_cache_nega)
noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
else:
noise_pred = noise_pred_posi
# Scheduler
latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
# Decode
self.load_models_to_device(['vae'])
frames = self.decode_video(latents, **tiler_kwargs)
@@ -1456,9 +1475,6 @@ class WanRepalceAnyoneVideoPipeline(BasePipeline):
frames=rearrange(frames[0], "C T H W -> T H W C")
return frames
def get_denoising_step_list(num_inference_steps, max_step=1000, min_step=0):
"""
根据步数自动生成 denoising_step_list,等间隔采样。
"""
return np.linspace(max_step, min_step, num_inference_steps, dtype=int).tolist()
step = (max_step - min_step) // num_inference_steps
return [max_step - step * i for i in range(num_inference_steps)]