Update sampler and print (#189)

This commit is contained in:
Bubbliiiing
2025-04-29 20:42:21 +08:00
committed by GitHub
parent 7d91e1369c
commit d11d6665bc
11 changed files with 90 additions and 19 deletions
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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.
+1 -1
View File
@@ -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.
+3 -3
View File
@@ -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
+19 -4
View File
@@ -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:
+4
View File
@@ -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}
+24 -3
View File
@@ -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:
+24 -3
View File
@@ -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:
+11 -1
View File
@@ -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
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