Update sampler and print (#189)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user