From d11d6665bc92f57b64e46a147e60dd737bce3625 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Tue, 29 Apr 2025 20:42:21 +0800 Subject: [PATCH] Update sampler and print (#189) --- examples/wan2.1/post_infer.py | 2 +- examples/wan2.1/post_infer_queue.py | 2 +- examples/wan2.1/post_infer_queue_i2v.py | 2 +- examples/wan2.1/predict_i2v.py | 2 +- examples/wan2.1/predict_t2v.py | 2 +- videox_fun/models/wan_transformer3d.py | 6 +++--- videox_fun/ui/cogvideox_fun_ui.py | 23 +++++++++++++++++---- videox_fun/ui/controller.py | 4 ++++ videox_fun/ui/wan_fun_ui.py | 27 ++++++++++++++++++++++--- videox_fun/ui/wan_ui.py | 27 ++++++++++++++++++++++--- videox_fun/utils/utils.py | 12 ++++++++++- 11 files changed, 90 insertions(+), 19 deletions(-) diff --git a/examples/wan2.1/post_infer.py b/examples/wan2.1/post_infer.py index 0b973b7..66390b9 100755 --- a/examples/wan2.1/post_infer.py +++ b/examples/wan2.1/post_infer.py @@ -99,7 +99,7 @@ if __name__ == '__main__': prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." negative_prompt_textbox = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion." # Sampler name - sampler_dropdown = "Flow" + sampler_dropdown = "Flow_Unipc" # Sampler steps sample_step_slider = 50 # height and width diff --git a/examples/wan2.1/post_infer_queue.py b/examples/wan2.1/post_infer_queue.py index 6ba1e66..2145f90 100755 --- a/examples/wan2.1/post_infer_queue.py +++ b/examples/wan2.1/post_infer_queue.py @@ -130,7 +130,7 @@ if __name__ == '__main__': prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # Sampler name - sampler_dropdown = "Flow" + sampler_dropdown = "Flow_Unipc" # Sampler steps sample_step_slider = 50 # height and width diff --git a/examples/wan2.1/post_infer_queue_i2v.py b/examples/wan2.1/post_infer_queue_i2v.py index 56cf5d0..659aab1 100755 --- a/examples/wan2.1/post_infer_queue_i2v.py +++ b/examples/wan2.1/post_infer_queue_i2v.py @@ -147,7 +147,7 @@ if __name__ == '__main__': prompt_textbox = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。" negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # Sampler name - sampler_dropdown = "Flow" + sampler_dropdown = "Flow_Unipc" # Sampler steps sample_step_slider = 50 # height and width diff --git a/examples/wan2.1/predict_i2v.py b/examples/wan2.1/predict_i2v.py index 0be52f4..7b7abde 100755 --- a/examples/wan2.1/predict_i2v.py +++ b/examples/wan2.1/predict_i2v.py @@ -78,7 +78,7 @@ config_path = "config/wan2.1/wan_civitai.yaml" model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-480P" # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" -sampler_name = "Flow" +sampler_name = "Flow_Unipc" # [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. # Used when the sampler is in "Flow_Unipc", "Flow_DPM++". # If you want to generate a 480p video, it is recommended to set the shift value to 3.0. diff --git a/examples/wan2.1/predict_t2v.py b/examples/wan2.1/predict_t2v.py index 62bc567..105b128 100755 --- a/examples/wan2.1/predict_t2v.py +++ b/examples/wan2.1/predict_t2v.py @@ -77,7 +77,7 @@ config_path = "config/wan2.1/wan_civitai.yaml" model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B" # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" -sampler_name = "Flow" +sampler_name = "Flow_Unipc" # [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. # Used when the sampler is in "Flow_Unipc", "Flow_DPM++". # If you want to generate a 480p video, it is recommended to set the shift value to 3.0. diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 9721554..77c3660 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -954,10 +954,10 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): e = self.time_embedding( sinusoidal_embedding_1d(self.freq_dim, t).float()) e0 = self.time_projection(e).unflatten(1, (6, self.dim)) - # to bfloat16 for saving memeory + # assert e.dtype == torch.float32 and e0.dtype == torch.float32 - e0 = e0.to(dtype) - e = e.to(dtype) + # e0 = e0.to(dtype) + # e = e.to(dtype) # context context_lens = None diff --git a/videox_fun/ui/cogvideox_fun_ui.py b/videox_fun/ui/cogvideox_fun_ui.py index f9e6f9b..f5441ad 100755 --- a/videox_fun/ui/cogvideox_fun_ui.py +++ b/videox_fun/ui/cogvideox_fun_ui.py @@ -17,7 +17,7 @@ from ..pipeline import (CogVideoXFunControlPipeline, CogVideoXFunInpaintPipeline, CogVideoXFunPipeline) from ..utils.fp8_optimization import convert_weight_dtype_wrapper from ..utils.lora_utils import merge_lora, unmerge_lora -from ..utils.utils import (get_image_to_video_latent, +from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, timer, get_video_to_video_latent, save_videos_grid) from .controller import (Fun_Controller, Fun_Controller_Client, all_cheduler_dict, css, ddpm_scheduler_dict, @@ -105,6 +105,7 @@ class CogVideoXFunController(Fun_Controller): print("Update diffusion transformer done") return gr.update() + @timer def generate( self, diffusion_transformer_dropdown, @@ -143,9 +144,11 @@ class CogVideoXFunController(Fun_Controller): ): self.clear_cache() + print(f"Input checking.") _, comment = self.input_check( resize_method, generation_method, start_image, end_image, validation_video,control_video, is_api ) + print(f"Input checking down") if comment != "OK": return "", comment is_image = True if generation_method == "Image Generation" else False @@ -156,21 +159,29 @@ class CogVideoXFunController(Fun_Controller): if self.lora_model_path != lora_model_dropdown: self.update_lora_model(lora_model_dropdown) + print(f"Load scheduler.") self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(self.pipeline.scheduler.config) + print(f"Load scheduler down.") if resize_method == "Resize according to Reference": + print(f"Calculate height and width according to Reference.") height_slider, width_slider = self.get_height_width_from_reference( base_resolution, start_image, validation_video, control_video, ) - if self.lora_model_path != "none": - # lora part - self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + if self.lora_model_path != "none": + print(f"Merge Lora.") + self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Merge Lora done.") + + print(f"Generate seed.") if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox)) else: seed_textbox = np.random.randint(0, 1e10) generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) + print(f"Generate seed done.") try: + print(f"Generation.") if self.model_type == "Inpaint": if self.transformer.config.in_channels != self.vae.config.latent_channels: if generation_method == "Long Video Generation": @@ -294,11 +305,15 @@ class CogVideoXFunController(Fun_Controller): self.clear_cache() # lora part if self.lora_model_path != "none": + print(f"Unmerge Lora.") self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Unmerge Lora done.") + print(f"Saving outputs.") save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=8 ) + print(f"Saving outputs done.") if is_image or length_slider == 1: if is_api: diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py index 951b498..6427ad2 100755 --- a/videox_fun/ui/controller.py +++ b/videox_fun/ui/controller.py @@ -24,6 +24,8 @@ from safetensors import safe_open from ..data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio from ..utils.utils import save_videos_grid +from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler +from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from ..dist import set_multi_gpus_devices gradio_version = pkg_resources.get_distribution("gradio").version @@ -49,6 +51,8 @@ ddpm_scheduler_dict = { } flow_scheduler_dict = { "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, } all_cheduler_dict = {**ddpm_scheduler_dict, **flow_scheduler_dict} diff --git a/videox_fun/ui/wan_fun_ui.py b/videox_fun/ui/wan_fun_ui.py index 96d97f6..fc82ed7 100755 --- a/videox_fun/ui/wan_fun_ui.py +++ b/videox_fun/ui/wan_fun_ui.py @@ -21,7 +21,7 @@ from ..utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper, replace_parameters_by_name) from ..utils.lora_utils import merge_lora, unmerge_lora -from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, +from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, timer, get_video_to_video_latent, save_videos_grid) from .controller import (Fun_Controller, Fun_Controller_Client, all_cheduler_dict, css, ddpm_scheduler_dict, @@ -138,6 +138,7 @@ class Wan_Fun_Controller(Fun_Controller): print("Update diffusion transformer done") return gr.update() + @timer def generate( self, diffusion_transformer_dropdown, @@ -176,9 +177,11 @@ class Wan_Fun_Controller(Fun_Controller): ): self.clear_cache() + print(f"Input checking.") _, comment = self.input_check( resize_method, generation_method, start_image, end_image, validation_video,control_video, is_api ) + print(f"Input checking down") if comment != "OK": return "", comment is_image = True if generation_method == "Image Generation" else False @@ -189,15 +192,23 @@ class Wan_Fun_Controller(Fun_Controller): if self.lora_model_path != lora_model_dropdown: self.update_lora_model(lora_model_dropdown) - self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(self.pipeline.scheduler.config) + print(f"Load scheduler.") + scheduler_config = self.pipeline.scheduler.config + if sampler_dropdown == "Flow_Unipc" or sampler_dropdown == "Flow_DPM++": + scheduler_config['scheduler_kwargs']['shift'] = 1 + self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(scheduler_config) + print(f"Load scheduler down.") if resize_method == "Resize according to Reference": + print(f"Calculate height and width according to Reference.") height_slider, width_slider = self.get_height_width_from_reference( base_resolution, start_image, validation_video, control_video, ) + if self.lora_model_path != "none": - # lora part + print(f"Merge Lora.") self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Merge Lora done.") coefficients = get_teacache_coefficients(self.diffusion_transformer_dropdown) if enable_teacache else None if coefficients is not None: @@ -206,17 +217,22 @@ class Wan_Fun_Controller(Fun_Controller): coefficients, sample_step_slider, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) else: + print(f"Disable TeaCache.") self.pipeline.transformer.disable_teacache() + print(f"Generate seed.") if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox)) else: seed_textbox = np.random.randint(0, 1e10) generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) + print(f"Generate seed done.") if enable_riflex: + print(f"Enable riflex") latent_frames = (int(length_slider) - 1) // self.vae.config.temporal_compression_ratio + 1 self.pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames if not is_image else 1) try: + print(f"Generation.") if self.model_type == "Inpaint": if self.transformer.config.in_channels != self.vae.config.latent_channels: if validation_video is not None: @@ -283,6 +299,7 @@ class Wan_Fun_Controller(Fun_Controller): clip_image = clip_image, cfg_skip_ratio = cfg_skip_ratio, ).videos + print(f"Generation done.") except Exception as e: self.clear_cache() print(f"Error. error information is {str(e)}") @@ -296,11 +313,15 @@ class Wan_Fun_Controller(Fun_Controller): self.clear_cache() # lora part if self.lora_model_path != "none": + print(f"Unmerge Lora.") self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Unmerge Lora done.") + print(f"Saving outputs.") save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=16 ) + print(f"Saving outputs done.") if is_image or length_slider == 1: if is_api: diff --git a/videox_fun/ui/wan_ui.py b/videox_fun/ui/wan_ui.py index 8f9a3cb..8597bb7 100755 --- a/videox_fun/ui/wan_ui.py +++ b/videox_fun/ui/wan_ui.py @@ -20,7 +20,7 @@ from ..utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper, replace_parameters_by_name) from ..utils.lora_utils import merge_lora, unmerge_lora -from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, +from ..utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, timer, get_video_to_video_latent, save_videos_grid) from .controller import (Fun_Controller, Fun_Controller_Client, all_cheduler_dict, css, ddpm_scheduler_dict, @@ -130,6 +130,7 @@ class Wan_Controller(Fun_Controller): print("Update diffusion transformer done") return gr.update() + @timer def generate( self, diffusion_transformer_dropdown, @@ -168,9 +169,11 @@ class Wan_Controller(Fun_Controller): ): self.clear_cache() + print(f"Input checking.") _, comment = self.input_check( resize_method, generation_method, start_image, end_image, validation_video,control_video, is_api ) + print(f"Input checking down") if comment != "OK": return "", comment is_image = True if generation_method == "Image Generation" else False @@ -181,15 +184,23 @@ class Wan_Controller(Fun_Controller): if self.lora_model_path != lora_model_dropdown: self.update_lora_model(lora_model_dropdown) - self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(self.pipeline.scheduler.config) + print(f"Load scheduler.") + scheduler_config = self.pipeline.scheduler.config + if sampler_dropdown == "Flow_Unipc" or sampler_dropdown == "Flow_DPM++": + scheduler_config['scheduler_kwargs']['shift'] = 1 + self.pipeline.scheduler = self.scheduler_dict[sampler_dropdown].from_config(scheduler_config) + print(f"Load scheduler down.") if resize_method == "Resize according to Reference": + print(f"Calculate height and width according to Reference.") height_slider, width_slider = self.get_height_width_from_reference( base_resolution, start_image, validation_video, control_video, ) + if self.lora_model_path != "none": - # lora part + print(f"Merge Lora.") self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Merge Lora done.") coefficients = get_teacache_coefficients(self.diffusion_transformer_dropdown) if enable_teacache else None if coefficients is not None: @@ -198,17 +209,22 @@ class Wan_Controller(Fun_Controller): coefficients, sample_step_slider, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload ) else: + print(f"Disable TeaCache.") self.pipeline.transformer.disable_teacache() + print(f"Generate seed.") if int(seed_textbox) != -1 and seed_textbox != "": torch.manual_seed(int(seed_textbox)) else: seed_textbox = np.random.randint(0, 1e10) generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) + print(f"Generate seed done.") if enable_riflex: + print(f"Enable riflex") latent_frames = (int(length_slider) - 1) // self.vae.config.temporal_compression_ratio + 1 self.pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames if not is_image else 1) try: + print(f"Generation.") if self.model_type == "Inpaint": if self.transformer.config.in_channels != self.vae.config.latent_channels: if validation_video is not None: @@ -275,6 +291,7 @@ class Wan_Controller(Fun_Controller): clip_image = clip_image, cfg_skip_ratio = cfg_skip_ratio, ).videos + print(f"Generation done.") except Exception as e: self.clear_cache() print(f"Error. error information is {str(e)}") @@ -288,11 +305,15 @@ class Wan_Controller(Fun_Controller): self.clear_cache() # lora part if self.lora_model_path != "none": + print(f"Unmerge Lora.") self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) + print(f"Unmerge Lora done.") + print(f"Saving outputs.") save_sample_path = self.save_outputs( is_image, length_slider, sample, fps=16 ) + print(f"Saving outputs done.") if is_image or length_slider == 1: if is_api: diff --git a/videox_fun/utils/utils.py b/videox_fun/utils/utils.py index 13bf02a..07b16dc 100755 --- a/videox_fun/utils/utils.py +++ b/videox_fun/utils/utils.py @@ -4,6 +4,7 @@ import imageio import inspect import numpy as np import torch +import time import torchvision import cv2 from einops import rearrange @@ -245,4 +246,13 @@ def get_image_latent(ref_image=None, sample_size=None): ref_image = torch.from_numpy(np.array(ref_image)) ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255 - return ref_image \ No newline at end of file + return ref_image + +def timer(func): + def wrapper(*args, **kwargs): + start_time = time.time() + result = func(*args, **kwargs) + end_time = time.time() + print(f"function {func.__name__} running for {end_time - start_time} seconds") + return result + return wrapper \ No newline at end of file