Update Wan2.2 UI
This commit is contained in:
@@ -0,0 +1,79 @@
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
import torch
|
||||
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
|
||||
from videox_fun.api.api import (infer_forward_api,
|
||||
update_diffusion_transformer_api)
|
||||
from videox_fun.ui.controller import flow_scheduler_dict
|
||||
from videox_fun.ui.wan2_2_ui import ui, ui_client, ui_host
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Choose the ui mode
|
||||
# "normal" refers to the standard UI, which allows users to click to switch models, change model types, and more.
|
||||
# "host" represents the hosting mode, where the model is loaded directly at startup and can be accessed via
|
||||
# the API to return generation results.
|
||||
# "client" represents the client mode, offering a simple UI that sends requests to a remote API for generation.
|
||||
ui_mode = "host"
|
||||
|
||||
# GPU memory mode, which can be choosen in [model_full_load, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
|
||||
# model_full_load means that the entire model will be moved to the GPU.
|
||||
#
|
||||
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
|
||||
#
|
||||
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
|
||||
# and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
#
|
||||
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
|
||||
# resulting in slower speeds but saving a large amount of GPU memory.
|
||||
GPU_memory_mode = "sequential_cpu_offload"
|
||||
# Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
compile_dit = False
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
|
||||
# Server ip
|
||||
server_name = "127.0.0.1"
|
||||
server_port = 7860
|
||||
|
||||
# Config path
|
||||
config_path = "config/wan2.2/wan_civitai_i2v.yaml"
|
||||
# Params below is used when ui_mode = "host"
|
||||
# Model path of the pretrained model
|
||||
model_name = "models/Diffusion_Transformer/Wan2.2-I2V-A14B"
|
||||
# "Inpaint" or "Control"
|
||||
model_type = "Inpaint"
|
||||
|
||||
if ui_mode == "host":
|
||||
demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype)
|
||||
elif ui_mode == "client":
|
||||
demo, controller = ui_client(flow_scheduler_dict, model_name)
|
||||
else:
|
||||
demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, compile_dit, weight_dtype)
|
||||
|
||||
def gr_launch():
|
||||
# launch gradio
|
||||
app, _, _ = demo.queue(status_update_rate=1).launch(
|
||||
server_name=server_name,
|
||||
server_port=server_port,
|
||||
prevent_thread_lock=True
|
||||
)
|
||||
|
||||
# launch api
|
||||
infer_forward_api(None, app, controller)
|
||||
update_diffusion_transformer_api(None, app, controller)
|
||||
|
||||
gr_launch()
|
||||
|
||||
# not close the python
|
||||
while True:
|
||||
time.sleep(5)
|
||||
@@ -0,0 +1,91 @@
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
import gradio as gr
|
||||
import ray
|
||||
import torch
|
||||
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
|
||||
from videox_fun.api.api_multi_nodes import (MultiNodesEngine,
|
||||
multi_nodes_infer_forward_api)
|
||||
from videox_fun.ui.controller import flow_scheduler_dict
|
||||
from videox_fun.ui.wan2_2_ui import Wan2_2_Controller
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='xDiT HTTP Service')
|
||||
parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
|
||||
parser.add_argument(
|
||||
'--gpu_memory_mode', type=str, default="model_full_load", help='''
|
||||
GPU memory mode, which can be choosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
|
||||
model_full_load means that the entire model will be moved to the GPU.
|
||||
|
||||
model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
|
||||
and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
|
||||
model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
|
||||
|
||||
model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
|
||||
and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
'''
|
||||
)
|
||||
parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
|
||||
parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
|
||||
parser.add_argument(
|
||||
'--compile_dit', action='store_true', help='''
|
||||
Enable compile dit.
|
||||
Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
'''
|
||||
)
|
||||
parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
|
||||
parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
|
||||
parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
|
||||
parser.add_argument('--server_name', type=str, default="0.0.0.0", help='Server IP address')
|
||||
parser.add_argument('--server_port', type=int, default=7860, help='Server Port')
|
||||
parser.add_argument('--config_path', type=str, default="config/wan2.2/wan_civitai_i2v.yaml", help='Path to config file')
|
||||
parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/Wan2.2-I2V-A14B", help='Model path')
|
||||
parser.add_argument('--model_type', type=str, default="Inpaint", help='Model type (Inpaint/Control)')
|
||||
parser.add_argument('--savedir_sample', type=str, default=None, help='The save directory for samples')
|
||||
args = parser.parse_args()
|
||||
|
||||
weight_dtype = torch.float32
|
||||
if args.weight_dtype == "bf16":
|
||||
weight_dtype = torch.bfloat16
|
||||
elif args.weight_dtype == "fp16":
|
||||
weight_dtype = torch.float16
|
||||
|
||||
engine = MultiNodesEngine(
|
||||
world_size=args.world_size, Controller=Wan2_2_Controller,
|
||||
GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=args.config_path,
|
||||
ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
|
||||
fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
|
||||
)
|
||||
|
||||
def gr_launch():
|
||||
# launch gradio
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("")
|
||||
app, _, _ = demo.queue(status_update_rate=1).launch(
|
||||
server_name=args.server_name,
|
||||
server_port=args.server_port,
|
||||
prevent_thread_lock=True
|
||||
)
|
||||
|
||||
# launch api
|
||||
multi_nodes_infer_forward_api(None, app, engine)
|
||||
|
||||
gr_launch()
|
||||
|
||||
# not close the python
|
||||
while True:
|
||||
time.sleep(5)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,169 @@
|
||||
import base64
|
||||
import json
|
||||
import time
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
|
||||
import requests
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def post_diffusion_transformer(diffusion_transformer_path, url='http://127.0.0.1:7860'):
|
||||
datas = json.dumps({
|
||||
"diffusion_transformer_path": diffusion_transformer_path
|
||||
})
|
||||
r = requests.post(f'{url}/videox_fun/update_diffusion_transformer', data=datas, timeout=1500)
|
||||
data = r.content.decode('utf-8')
|
||||
return data
|
||||
|
||||
def post_update_edition(edition, url='http://0.0.0.0:7860'):
|
||||
datas = json.dumps({
|
||||
"edition": edition
|
||||
})
|
||||
r = requests.post(f'{url}/videox_fun/update_edition', data=datas, timeout=1500)
|
||||
data = r.content.decode('utf-8')
|
||||
return data
|
||||
|
||||
|
||||
def post_infer(
|
||||
generation_method,
|
||||
length_slider,
|
||||
url='http://127.0.0.1:7860',
|
||||
POST_TOKEN="",
|
||||
timeout=5000,
|
||||
base_model_path="none",
|
||||
lora_model_path="none",
|
||||
lora_alpha_slider=0.55,
|
||||
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_dropdown="Flow",
|
||||
sample_step_slider=50,
|
||||
width_slider=672,
|
||||
height_slider=384,
|
||||
cfg_scale_slider=6,
|
||||
seed_textbox=43,
|
||||
start_image = None
|
||||
):
|
||||
if start_image:
|
||||
try:
|
||||
if not start_image.startswith("http"):
|
||||
image = Image.open(start_image).convert("RGB")
|
||||
# 将图片转换为 Base64 编码
|
||||
buffered = BytesIO()
|
||||
image.save(buffered, format="JPEG")
|
||||
start_image = base64.b64encode(buffered.getvalue()).decode('utf-8')
|
||||
except Exception as e:
|
||||
print(f"Error processing start_image: {e}")
|
||||
raise
|
||||
|
||||
# Prepare the data payload
|
||||
datas = json.dumps({
|
||||
"base_model_path": base_model_path,
|
||||
"lora_model_path": lora_model_path,
|
||||
"lora_alpha_slider": lora_alpha_slider,
|
||||
"prompt_textbox": prompt_textbox,
|
||||
"negative_prompt_textbox": negative_prompt_textbox,
|
||||
"sampler_dropdown": sampler_dropdown,
|
||||
"sample_step_slider": sample_step_slider,
|
||||
"width_slider": width_slider,
|
||||
"height_slider": height_slider,
|
||||
"generation_method": generation_method,
|
||||
"length_slider": length_slider,
|
||||
"cfg_scale_slider": cfg_scale_slider,
|
||||
"seed_textbox": seed_textbox,
|
||||
|
||||
"start_image": start_image
|
||||
})
|
||||
|
||||
# Initialize session and set headers
|
||||
session = requests.session()
|
||||
session.headers.update({"Authorization": POST_TOKEN})
|
||||
|
||||
# Send POST request
|
||||
if url[-1] == "/":
|
||||
url = url[:-1]
|
||||
post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout)
|
||||
|
||||
data = post_r.content.decode('utf-8')
|
||||
return data
|
||||
|
||||
if __name__ == '__main__':
|
||||
# initiate time
|
||||
time_start = time.time()
|
||||
|
||||
# The Url you want to post
|
||||
POST_URL = 'http://0.0.0.0:7860'
|
||||
# Used in EAS. If you don't need Authorization, please set it to empty string.
|
||||
TOKEN = ''
|
||||
|
||||
# -------------------------- #
|
||||
# Step 1: update edition
|
||||
# -------------------------- #
|
||||
# diffusion_transformer_path = "models/Diffusion_Transformer/Wan2.2-I2V-A14B"
|
||||
# outputs = post_diffusion_transformer(diffusion_transformer_path)
|
||||
# print('Output update edition: ', outputs)
|
||||
|
||||
# -------------------------- #
|
||||
# Step 2: infer
|
||||
# -------------------------- #
|
||||
# "Video Generation" and "Image Generation"
|
||||
generation_method = "Video Generation"
|
||||
# Video length
|
||||
length_slider = 49
|
||||
# Used in Lora models
|
||||
lora_model_path = "none"
|
||||
lora_alpha_slider = 0.55
|
||||
# Prompts
|
||||
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 steps
|
||||
sample_step_slider = 50
|
||||
# height and width
|
||||
width_slider = 832
|
||||
height_slider = 480
|
||||
# cfg scale
|
||||
cfg_scale_slider = 6
|
||||
seed_textbox = 43
|
||||
# 起始图片路径
|
||||
start_image_path = "asset/3.png" # 替换为实际的图片路径
|
||||
|
||||
outputs = post_infer(
|
||||
generation_method,
|
||||
length_slider,
|
||||
lora_model_path=lora_model_path,
|
||||
lora_alpha_slider=lora_alpha_slider,
|
||||
prompt_textbox=prompt_textbox,
|
||||
negative_prompt_textbox=negative_prompt_textbox,
|
||||
sampler_dropdown=sampler_dropdown,
|
||||
sample_step_slider=sample_step_slider,
|
||||
width_slider=width_slider,
|
||||
height_slider=height_slider,
|
||||
cfg_scale_slider=cfg_scale_slider,
|
||||
seed_textbox=seed_textbox,
|
||||
url=POST_URL,
|
||||
POST_TOKEN=TOKEN,
|
||||
start_image=start_image_path
|
||||
)
|
||||
|
||||
# Get decoded data
|
||||
outputs = json.loads(outputs)
|
||||
base64_encoding = outputs["base64_encoding"]
|
||||
decoded_data = base64.b64decode(base64_encoding)
|
||||
|
||||
is_image = True if generation_method == "Image Generation" else False
|
||||
if is_image or length_slider == 1:
|
||||
file_path = "1.png"
|
||||
else:
|
||||
file_path = "1.mp4"
|
||||
with open(file_path, "wb") as file:
|
||||
file.write(decoded_data)
|
||||
|
||||
# End of record time
|
||||
# The calculated time difference is the execution time of the program, expressed in seconds / s
|
||||
time_end = time.time()
|
||||
time_sum = (time_end - time_start)
|
||||
print('# --------------------------------------------------------- #')
|
||||
print(f'# Total expenditure: {time_sum}s')
|
||||
print('# --------------------------------------------------------- #')
|
||||
@@ -747,6 +747,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
in_dim_control_adapter=24,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
cross_attn_type=None,
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -786,7 +787,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
|
||||
super().__init__()
|
||||
|
||||
assert model_type in ['t2v', 'i2v']
|
||||
assert model_type in ['t2v', 'i2v', 'ti2v']
|
||||
self.model_type = model_type
|
||||
|
||||
self.patch_size = patch_size
|
||||
@@ -816,7 +817,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
|
||||
# blocks
|
||||
cross_attn_type = 't2v_cross_attn' if model_type == 't2v' else 'i2v_cross_attn'
|
||||
if cross_attn_type is None:
|
||||
cross_attn_type = 't2v_cross_attn' if model_type == 't2v' else 'i2v_cross_attn'
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(cross_attn_type, dim, ffn_dim, num_heads,
|
||||
window_size, qk_norm, cross_attn_norm, eps)
|
||||
@@ -1404,7 +1406,6 @@ class Wan2_2Transformer3DModel(WanTransformer3DModel):
|
||||
# _no_split_modules = ['WanAttentionBlock']
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
model_type='t2v',
|
||||
@@ -1463,66 +1464,30 @@ class Wan2_2Transformer3DModel(WanTransformer3DModel):
|
||||
eps (`float`, *optional*, defaults to 1e-6):
|
||||
Epsilon value for normalization layers
|
||||
"""
|
||||
super().__init__()
|
||||
assert model_type in ['t2v', 'i2v', 'ti2v']
|
||||
self.model_type = model_type
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.freq_dim = freq_dim
|
||||
self.text_dim = text_dim
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.text_embedding = nn.Sequential(
|
||||
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(dim, dim))
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
# blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock("cross_attn", dim, ffn_dim, num_heads, window_size, qk_norm,
|
||||
cross_attn_norm, eps) for _ in range(num_layers)
|
||||
])
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
block.self_attn.layer_idx = layer_idx
|
||||
block.self_attn.num_layers = self.num_layers
|
||||
|
||||
# head
|
||||
self.head = Head(dim, out_dim, patch_size, eps)
|
||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
|
||||
if add_control_adapter:
|
||||
self.control_adapter = SimpleAdapter(in_dim_control_adapter, dim, kernel_size=patch_size[1:], stride=patch_size[1:])
|
||||
else:
|
||||
self.control_adapter = None
|
||||
|
||||
if add_ref_conv:
|
||||
self.ref_conv = nn.Conv2d(in_dim_ref_conv, dim, kernel_size=patch_size[1:], stride=patch_size[1:])
|
||||
else:
|
||||
self.ref_conv = None
|
||||
super().__init__(
|
||||
model_type=model_type,
|
||||
patch_size=patch_size,
|
||||
text_len=text_len,
|
||||
in_dim=in_dim,
|
||||
dim=dim,
|
||||
ffn_dim=ffn_dim,
|
||||
freq_dim=freq_dim,
|
||||
text_dim=text_dim,
|
||||
out_dim=out_dim,
|
||||
num_heads=num_heads,
|
||||
num_layers=num_layers,
|
||||
window_size=window_size,
|
||||
qk_norm=qk_norm,
|
||||
cross_attn_norm=cross_attn_norm,
|
||||
eps=eps,
|
||||
in_channels=in_channels,
|
||||
hidden_size=hidden_size,
|
||||
add_control_adapter=add_control_adapter,
|
||||
in_dim_control_adapter=in_dim_control_adapter,
|
||||
add_ref_conv=add_ref_conv,
|
||||
in_dim_ref_conv=in_dim_ref_conv,
|
||||
cross_attn_type="cross_attn"
|
||||
)
|
||||
|
||||
if hasattr(self, "img_emb"):
|
||||
del self.img_emb
|
||||
|
||||
# initialize weights
|
||||
self.init_weights()
|
||||
del self.img_emb
|
||||
@@ -0,0 +1,740 @@
|
||||
"""Modified from https://github.com/guoyww/AnimateDiff/blob/main/app.py
|
||||
"""
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from PIL import Image
|
||||
from safetensors import safe_open
|
||||
|
||||
from ..data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio
|
||||
from ..models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
|
||||
WanT5EncoderModel, Wan2_2Transformer3DModel)
|
||||
from ..models.cache_utils import get_teacache_coefficients
|
||||
from ..pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline
|
||||
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, timer,
|
||||
get_video_to_video_latent, save_videos_grid)
|
||||
from .controller import (Fun_Controller, Fun_Controller_Client,
|
||||
all_cheduler_dict, css, ddpm_scheduler_dict,
|
||||
flow_scheduler_dict, gradio_version,
|
||||
gradio_version_is_above_4)
|
||||
from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k,
|
||||
create_cfg_skip_params,
|
||||
create_fake_finetune_models_checkpoints,
|
||||
create_fake_height_width, create_fake_model_checkpoints,
|
||||
create_fake_model_type, create_finetune_models_checkpoints,
|
||||
create_generation_method,
|
||||
create_generation_methods_and_video_length,
|
||||
create_height_width, create_model_checkpoints,
|
||||
create_model_type, create_prompts, create_samplers,
|
||||
create_teacache_params, create_ui_outputs)
|
||||
from ..dist import set_multi_gpus_devices, shard_model
|
||||
|
||||
|
||||
class Wan2_2_Controller(Fun_Controller):
|
||||
def update_diffusion_transformer(self, diffusion_transformer_dropdown):
|
||||
print(f"Update diffusion transformer: {diffusion_transformer_dropdown}")
|
||||
self.model_name = diffusion_transformer_dropdown
|
||||
self.diffusion_transformer_dropdown = diffusion_transformer_dropdown
|
||||
if diffusion_transformer_dropdown == "none":
|
||||
return gr.update()
|
||||
self.vae = AutoencoderKLWan.from_pretrained(
|
||||
os.path.join(diffusion_transformer_dropdown, self.config['vae_kwargs'].get('vae_subpath', 'vae')),
|
||||
additional_kwargs=OmegaConf.to_container(self.config['vae_kwargs']),
|
||||
).to(self.weight_dtype)
|
||||
|
||||
# Get Transformer
|
||||
self.transformer = Wan2_2Transformer3DModel.from_pretrained(
|
||||
os.path.join(diffusion_transformer_dropdown, self.config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')),
|
||||
transformer_additional_kwargs=OmegaConf.to_container(self.config['transformer_additional_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=self.weight_dtype,
|
||||
)
|
||||
self.transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
|
||||
os.path.join(diffusion_transformer_dropdown, self.config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
|
||||
transformer_additional_kwargs=OmegaConf.to_container(self.config['transformer_additional_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=self.weight_dtype,
|
||||
)
|
||||
|
||||
# Get Tokenizer
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
os.path.join(diffusion_transformer_dropdown, self.config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
|
||||
)
|
||||
|
||||
# Get Text encoder
|
||||
self.text_encoder = WanT5EncoderModel.from_pretrained(
|
||||
os.path.join(diffusion_transformer_dropdown, self.config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
|
||||
additional_kwargs=OmegaConf.to_container(self.config['text_encoder_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=self.weight_dtype,
|
||||
)
|
||||
self.text_encoder = self.text_encoder.eval()
|
||||
|
||||
Choosen_Scheduler = self.scheduler_dict[list(self.scheduler_dict.keys())[0]]
|
||||
self.scheduler = Choosen_Scheduler(
|
||||
**filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(self.config['scheduler_kwargs']))
|
||||
)
|
||||
|
||||
# Get pipeline
|
||||
if self.model_type == "Inpaint":
|
||||
if self.transformer.config.in_channels != self.vae.config.latent_channels:
|
||||
self.pipeline = Wan2_2I2VPipeline(
|
||||
vae=self.vae,
|
||||
tokenizer=self.tokenizer,
|
||||
text_encoder=self.text_encoder,
|
||||
transformer=self.transformer,
|
||||
transformer_2=self.transformer_2,
|
||||
scheduler=self.scheduler,
|
||||
)
|
||||
else:
|
||||
self.pipeline = Wan2_2Pipeline(
|
||||
vae=self.vae,
|
||||
tokenizer=self.tokenizer,
|
||||
text_encoder=self.text_encoder,
|
||||
transformer=self.transformer,
|
||||
transformer_2=self.transformer_2,
|
||||
scheduler=self.scheduler,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Not support now")
|
||||
|
||||
if self.ulysses_degree > 1 or self.ring_degree > 1:
|
||||
from functools import partial
|
||||
self.transformer.enable_multi_gpus_inference()
|
||||
if self.fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
|
||||
self.pipeline.transformer = shard_fn(self.pipeline.transformer)
|
||||
self.pipeline.transformer_2 = shard_fn(self.pipeline.transformer_2)
|
||||
print("Add FSDP DIT")
|
||||
if self.fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=self.device, param_dtype=self.weight_dtype)
|
||||
self.pipeline.text_encoder = shard_fn(self.pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if self.compile_dit:
|
||||
for i in range(len(self.pipeline.transformer.blocks)):
|
||||
self.pipeline.transformer.blocks[i] = torch.compile(self.pipeline.transformer.blocks[i])
|
||||
for i in range(len(self.pipeline.transformer_2.blocks)):
|
||||
self.pipeline.transformer_2.blocks[i] = torch.compile(self.pipeline.transformer_2.blocks[i])
|
||||
print("Add Compile")
|
||||
|
||||
if self.GPU_memory_mode == "sequential_cpu_offload":
|
||||
replace_parameters_by_name(self.transformer, ["modulation",], device=self.device)
|
||||
replace_parameters_by_name(self.transformer_2, ["modulation",], device=self.device)
|
||||
self.transformer.freqs = self.transformer.freqs.to(device=self.device)
|
||||
self.transformer_2.freqs = self.transformer_2.freqs.to(device=self.device)
|
||||
self.pipeline.enable_sequential_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_model_weight_to_float8(self.transformer_2, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
|
||||
convert_weight_dtype_wrapper(self.transformer_2, self.weight_dtype)
|
||||
self.pipeline.enable_model_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_cpu_offload":
|
||||
self.pipeline.enable_model_cpu_offload(device=self.device)
|
||||
elif self.GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(self.transformer, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_model_weight_to_float8(self.transformer_2, exclude_module_name=["modulation",], device=self.device)
|
||||
convert_weight_dtype_wrapper(self.transformer, self.weight_dtype)
|
||||
convert_weight_dtype_wrapper(self.transformer_2, self.weight_dtype)
|
||||
self.pipeline.to(self.device)
|
||||
else:
|
||||
self.pipeline.to(self.device)
|
||||
print("Update diffusion transformer done")
|
||||
return gr.update()
|
||||
|
||||
@timer
|
||||
def generate(
|
||||
self,
|
||||
diffusion_transformer_dropdown,
|
||||
base_model_dropdown,
|
||||
lora_model_dropdown,
|
||||
lora_alpha_slider,
|
||||
prompt_textbox,
|
||||
negative_prompt_textbox,
|
||||
sampler_dropdown,
|
||||
sample_step_slider,
|
||||
resize_method,
|
||||
width_slider,
|
||||
height_slider,
|
||||
base_resolution,
|
||||
generation_method,
|
||||
length_slider,
|
||||
overlap_video_length,
|
||||
partial_video_length,
|
||||
cfg_scale_slider,
|
||||
start_image,
|
||||
end_image,
|
||||
validation_video,
|
||||
validation_video_mask,
|
||||
control_video,
|
||||
denoise_strength,
|
||||
seed_textbox,
|
||||
ref_image = None,
|
||||
enable_teacache = None,
|
||||
teacache_threshold = None,
|
||||
num_skip_start_steps = None,
|
||||
teacache_offload = None,
|
||||
cfg_skip_ratio = None,
|
||||
enable_riflex = None,
|
||||
riflex_k = None,
|
||||
fps = None,
|
||||
is_api = False,
|
||||
):
|
||||
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
|
||||
|
||||
if self.base_model_path != base_model_dropdown:
|
||||
self.update_base_model(base_model_dropdown)
|
||||
|
||||
if self.lora_model_path != lora_model_dropdown:
|
||||
self.update_lora_model(lora_model_dropdown)
|
||||
|
||||
print(f"Load scheduler.")
|
||||
scheduler_config = self.pipeline.scheduler.config
|
||||
if sampler_dropdown == "Flow_Unipc" or sampler_dropdown == "Flow_DPM++":
|
||||
scheduler_config['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":
|
||||
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:
|
||||
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
|
||||
self.pipeline.transformer.enable_teacache(
|
||||
coefficients, sample_step_slider, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
|
||||
)
|
||||
self.pipeline.transformer_2.share_teacache(
|
||||
self.pipeline.transformer
|
||||
)
|
||||
else:
|
||||
print(f"Disable TeaCache.")
|
||||
self.pipeline.transformer.disable_teacache()
|
||||
self.pipeline.transformer_2.disable_teacache()
|
||||
|
||||
if cfg_skip_ratio is not None and cfg_skip_ratio >= 0:
|
||||
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
|
||||
self.pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, sample_step_slider)
|
||||
self.pipeline.transformer_2.share_cfg_skip(self.pipeline.transformer)
|
||||
|
||||
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 fps is None:
|
||||
fps = 16
|
||||
boundary = self.config['transformer_additional_kwargs'].get('boundary', 0.875)
|
||||
|
||||
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)
|
||||
self.pipeline.transformer_2.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:
|
||||
input_video, input_video_mask, _, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=fps)
|
||||
else:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
|
||||
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
boundary = boundary
|
||||
).videos
|
||||
else:
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
boundary = boundary
|
||||
).videos
|
||||
else:
|
||||
if ref_image is not None:
|
||||
ref_image = get_image_latent(ref_image, sample_size=(height_slider, width_slider))
|
||||
|
||||
if start_image is not None:
|
||||
start_image = get_image_latent(start_image, sample_size=(height_slider, width_slider))
|
||||
|
||||
input_video, input_video_mask, _, _ = get_video_to_video_latent(control_video, video_length=length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=fps, ref_image=None)
|
||||
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
|
||||
control_video = input_video,
|
||||
ref_image = ref_image,
|
||||
start_image = start_image,
|
||||
boundary = boundary
|
||||
).videos
|
||||
print(f"Generation done.")
|
||||
except Exception as e:
|
||||
self.auto_model_clear_cache(self.pipeline.transformer)
|
||||
self.auto_model_clear_cache(self.pipeline.text_encoder)
|
||||
self.auto_model_clear_cache(self.pipeline.vae)
|
||||
self.clear_cache()
|
||||
|
||||
print(f"Error. error information is {str(e)}")
|
||||
if self.lora_model_path != "none":
|
||||
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
|
||||
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2")
|
||||
if is_api:
|
||||
return "", f"Error. error information is {str(e)}"
|
||||
else:
|
||||
return gr.update(), gr.update(), f"Error. error information is {str(e)}"
|
||||
|
||||
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)
|
||||
self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2")
|
||||
print(f"Unmerge Lora done.")
|
||||
|
||||
print(f"Saving outputs.")
|
||||
save_sample_path = self.save_outputs(
|
||||
is_image, length_slider, sample, fps=fps
|
||||
)
|
||||
print(f"Saving outputs done.")
|
||||
|
||||
if is_image or length_slider == 1:
|
||||
if is_api:
|
||||
return save_sample_path, "Success"
|
||||
else:
|
||||
if gradio_version_is_above_4:
|
||||
return gr.Image(value=save_sample_path, visible=True), gr.Video(value=None, visible=False), "Success"
|
||||
else:
|
||||
return gr.Image.update(value=save_sample_path, visible=True), gr.Video.update(value=None, visible=False), "Success"
|
||||
else:
|
||||
if is_api:
|
||||
return save_sample_path, "Success"
|
||||
else:
|
||||
if gradio_version_is_above_4:
|
||||
return gr.Image(visible=False, value=None), gr.Video(value=save_sample_path, visible=True), "Success"
|
||||
else:
|
||||
return gr.Image.update(visible=False, value=None), gr.Video.update(value=save_sample_path, visible=True), "Success"
|
||||
|
||||
Wan2_2_Controller_Host = Wan2_2_Controller
|
||||
Wan2_2_Controller_Client = Fun_Controller_Client
|
||||
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, savedir_sample=None):
|
||||
controller = Wan2_2_Controller(
|
||||
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=config_path, compile_dit=compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""
|
||||
# Wan2_2:
|
||||
"""
|
||||
)
|
||||
with gr.Column(variant="panel"):
|
||||
model_type = create_model_type(visible=False)
|
||||
diffusion_transformer_dropdown, diffusion_transformer_refresh_button = \
|
||||
create_model_checkpoints(controller, visible=True)
|
||||
base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button = \
|
||||
create_finetune_models_checkpoints(controller, visible=True)
|
||||
|
||||
with gr.Row():
|
||||
enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \
|
||||
create_teacache_params(True, 0.10, 1, False)
|
||||
cfg_skip_ratio = create_cfg_skip_params(0)
|
||||
enable_riflex, riflex_k = create_cfg_riflex_k(False, 6)
|
||||
|
||||
with gr.Column(variant="panel"):
|
||||
prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走")
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
sampler_dropdown, sample_step_slider = create_samplers(controller)
|
||||
|
||||
resize_method, width_slider, height_slider, base_resolution = create_height_width(
|
||||
default_height = 480, default_width = 832, maximum_height = 1344,
|
||||
maximum_width = 1344,
|
||||
)
|
||||
generation_method, length_slider, overlap_video_length, partial_video_length = \
|
||||
create_generation_methods_and_video_length(
|
||||
["Video Generation", "Image Generation"],
|
||||
default_video_length=81,
|
||||
maximum_video_length=161,
|
||||
)
|
||||
image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
|
||||
["Text to Video (文本到视频)", "Image to Video (图片到视频)"], prompt_textbox, support_end_image=False
|
||||
)
|
||||
cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4)
|
||||
|
||||
generate_button = gr.Button(value="Generate (生成)", variant='primary')
|
||||
|
||||
result_image, result_video, infer_progress = create_ui_outputs()
|
||||
|
||||
model_type.change(
|
||||
fn=controller.update_model_type,
|
||||
inputs=[model_type],
|
||||
outputs=[]
|
||||
)
|
||||
|
||||
def upload_generation_method(generation_method):
|
||||
if generation_method == "Video Generation":
|
||||
return [gr.update(visible=True, maximum=161, value=81, interactive=True), gr.update(visible=False), gr.update(visible=False)]
|
||||
elif generation_method == "Image Generation":
|
||||
return [gr.update(minimum=1, maximum=1, value=1, interactive=False), gr.update(visible=False), gr.update(visible=False)]
|
||||
else:
|
||||
return [gr.update(visible=True, maximum=1344), gr.update(visible=True), gr.update(visible=True)]
|
||||
generation_method.change(
|
||||
upload_generation_method, generation_method, [length_slider, overlap_video_length, partial_video_length]
|
||||
)
|
||||
|
||||
def upload_source_method(source_method):
|
||||
if source_method == "Text to Video (文本到视频)":
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)]
|
||||
elif source_method == "Image to Video (图片到视频)":
|
||||
return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)]
|
||||
elif source_method == "Video to Video (视频到视频)":
|
||||
return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()]
|
||||
source_method.change(
|
||||
upload_source_method, source_method, [
|
||||
image_to_video_col, video_to_video_col, control_video_col, start_image, end_image,
|
||||
validation_video, validation_video_mask, control_video
|
||||
]
|
||||
)
|
||||
|
||||
def upload_resize_method(resize_method):
|
||||
if resize_method == "Generate by":
|
||||
return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)]
|
||||
resize_method.change(
|
||||
upload_resize_method, resize_method, [width_slider, height_slider, base_resolution]
|
||||
)
|
||||
|
||||
generate_button.click(
|
||||
fn=controller.generate,
|
||||
inputs=[
|
||||
diffusion_transformer_dropdown,
|
||||
base_model_dropdown,
|
||||
lora_model_dropdown,
|
||||
lora_alpha_slider,
|
||||
prompt_textbox,
|
||||
negative_prompt_textbox,
|
||||
sampler_dropdown,
|
||||
sample_step_slider,
|
||||
resize_method,
|
||||
width_slider,
|
||||
height_slider,
|
||||
base_resolution,
|
||||
generation_method,
|
||||
length_slider,
|
||||
overlap_video_length,
|
||||
partial_video_length,
|
||||
cfg_scale_slider,
|
||||
start_image,
|
||||
end_image,
|
||||
validation_video,
|
||||
validation_video_mask,
|
||||
control_video,
|
||||
denoise_strength,
|
||||
seed_textbox,
|
||||
ref_image,
|
||||
enable_teacache,
|
||||
teacache_threshold,
|
||||
num_skip_start_steps,
|
||||
teacache_offload,
|
||||
cfg_skip_ratio,
|
||||
enable_riflex,
|
||||
riflex_k,
|
||||
],
|
||||
outputs=[result_image, result_video, infer_progress]
|
||||
)
|
||||
return demo, controller
|
||||
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, compile_dit, weight_dtype, savedir_sample=None):
|
||||
controller = Wan2_2_Controller_Host(
|
||||
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
|
||||
config_path=config_path, compile_dit=compile_dit,
|
||||
weight_dtype=weight_dtype, savedir_sample=savedir_sample,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""
|
||||
# Wan2_2:
|
||||
"""
|
||||
)
|
||||
with gr.Column(variant="panel"):
|
||||
model_type = create_fake_model_type(visible=False)
|
||||
diffusion_transformer_dropdown = create_fake_model_checkpoints(model_name, visible=True)
|
||||
base_model_dropdown, lora_model_dropdown, lora_alpha_slider = create_fake_finetune_models_checkpoints(visible=True)
|
||||
|
||||
with gr.Row():
|
||||
enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \
|
||||
create_teacache_params(True, 0.10, 1, False)
|
||||
cfg_skip_ratio = create_cfg_skip_params(0)
|
||||
enable_riflex, riflex_k = create_cfg_riflex_k(False, 6)
|
||||
|
||||
with gr.Column(variant="panel"):
|
||||
prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走")
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
sampler_dropdown, sample_step_slider = create_samplers(controller)
|
||||
|
||||
resize_method, width_slider, height_slider, base_resolution = create_height_width(
|
||||
default_height = 480, default_width = 832, maximum_height = 1344,
|
||||
maximum_width = 1344,
|
||||
)
|
||||
generation_method, length_slider, overlap_video_length, partial_video_length = \
|
||||
create_generation_methods_and_video_length(
|
||||
["Video Generation", "Image Generation"],
|
||||
default_video_length=81,
|
||||
maximum_video_length=161,
|
||||
)
|
||||
image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
|
||||
["Text to Video (文本到视频)", "Image to Video (图片到视频)"], prompt_textbox
|
||||
)
|
||||
cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4)
|
||||
|
||||
generate_button = gr.Button(value="Generate (生成)", variant='primary')
|
||||
|
||||
result_image, result_video, infer_progress = create_ui_outputs()
|
||||
|
||||
def upload_generation_method(generation_method):
|
||||
if generation_method == "Video Generation":
|
||||
return gr.update(visible=True, minimum=1, maximum=161, value=81, interactive=True)
|
||||
elif generation_method == "Image Generation":
|
||||
return gr.update(minimum=1, maximum=1, value=1, interactive=False)
|
||||
generation_method.change(
|
||||
upload_generation_method, generation_method, [length_slider]
|
||||
)
|
||||
|
||||
def upload_source_method(source_method):
|
||||
if source_method == "Text to Video (文本到视频)":
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)]
|
||||
elif source_method == "Image to Video (图片到视频)":
|
||||
return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)]
|
||||
elif source_method == "Video to Video (视频到视频)":
|
||||
return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()]
|
||||
source_method.change(
|
||||
upload_source_method, source_method, [
|
||||
image_to_video_col, video_to_video_col, control_video_col, start_image, end_image,
|
||||
validation_video, validation_video_mask, control_video
|
||||
]
|
||||
)
|
||||
|
||||
def upload_resize_method(resize_method):
|
||||
if resize_method == "Generate by":
|
||||
return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)]
|
||||
resize_method.change(
|
||||
upload_resize_method, resize_method, [width_slider, height_slider, base_resolution]
|
||||
)
|
||||
|
||||
generate_button.click(
|
||||
fn=controller.generate,
|
||||
inputs=[
|
||||
diffusion_transformer_dropdown,
|
||||
base_model_dropdown,
|
||||
lora_model_dropdown,
|
||||
lora_alpha_slider,
|
||||
prompt_textbox,
|
||||
negative_prompt_textbox,
|
||||
sampler_dropdown,
|
||||
sample_step_slider,
|
||||
resize_method,
|
||||
width_slider,
|
||||
height_slider,
|
||||
base_resolution,
|
||||
generation_method,
|
||||
length_slider,
|
||||
overlap_video_length,
|
||||
partial_video_length,
|
||||
cfg_scale_slider,
|
||||
start_image,
|
||||
end_image,
|
||||
validation_video,
|
||||
validation_video_mask,
|
||||
control_video,
|
||||
denoise_strength,
|
||||
seed_textbox,
|
||||
ref_image,
|
||||
enable_teacache,
|
||||
teacache_threshold,
|
||||
num_skip_start_steps,
|
||||
teacache_offload,
|
||||
cfg_skip_ratio,
|
||||
enable_riflex,
|
||||
riflex_k,
|
||||
],
|
||||
outputs=[result_image, result_video, infer_progress]
|
||||
)
|
||||
return demo, controller
|
||||
|
||||
def ui_client(scheduler_dict, model_name, savedir_sample=None):
|
||||
controller = Wan2_2_Controller_Client(scheduler_dict, savedir_sample)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""
|
||||
# Wan2_2:
|
||||
"""
|
||||
)
|
||||
with gr.Column(variant="panel"):
|
||||
diffusion_transformer_dropdown = create_fake_model_checkpoints(model_name, visible=True)
|
||||
base_model_dropdown, lora_model_dropdown, lora_alpha_slider = create_fake_finetune_models_checkpoints(visible=True)
|
||||
|
||||
with gr.Row():
|
||||
enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload = \
|
||||
create_teacache_params(True, 0.10, 1, False)
|
||||
cfg_skip_ratio = create_cfg_skip_params(0)
|
||||
enable_riflex, riflex_k = create_cfg_riflex_k(False, 6)
|
||||
|
||||
with gr.Column(variant="panel"):
|
||||
prompt_textbox, negative_prompt_textbox = create_prompts(negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走")
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
sampler_dropdown, sample_step_slider = create_samplers(controller, maximum_step=50)
|
||||
|
||||
resize_method, width_slider, height_slider, base_resolution = create_fake_height_width(
|
||||
default_height = 480, default_width = 832, maximum_height = 1344,
|
||||
maximum_width = 1344,
|
||||
)
|
||||
generation_method, length_slider, overlap_video_length, partial_video_length = \
|
||||
create_generation_methods_and_video_length(
|
||||
["Video Generation", "Image Generation"],
|
||||
default_video_length=81,
|
||||
maximum_video_length=161,
|
||||
)
|
||||
image_to_video_col, video_to_video_col, control_video_col, source_method, start_image, template_gallery, end_image, validation_video, validation_video_mask, denoise_strength, control_video, ref_image = create_generation_method(
|
||||
["Text to Video (文本到视频)", "Image to Video (图片到视频)"], prompt_textbox
|
||||
)
|
||||
|
||||
cfg_scale_slider, seed_textbox, seed_button = create_cfg_and_seedbox(gradio_version_is_above_4)
|
||||
|
||||
generate_button = gr.Button(value="Generate (生成)", variant='primary')
|
||||
|
||||
result_image, result_video, infer_progress = create_ui_outputs()
|
||||
|
||||
def upload_generation_method(generation_method):
|
||||
if generation_method == "Video Generation":
|
||||
return gr.update(visible=True, minimum=5, maximum=161, value=49, interactive=True)
|
||||
elif generation_method == "Image Generation":
|
||||
return gr.update(minimum=1, maximum=1, value=1, interactive=False)
|
||||
generation_method.change(
|
||||
upload_generation_method, generation_method, [length_slider]
|
||||
)
|
||||
|
||||
def upload_source_method(source_method):
|
||||
if source_method == "Text to Video (文本到视频)":
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)]
|
||||
elif source_method == "Image to Video (图片到视频)":
|
||||
return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(), gr.update()]
|
||||
source_method.change(
|
||||
upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video, validation_video_mask]
|
||||
)
|
||||
|
||||
def upload_resize_method(resize_method):
|
||||
if resize_method == "Generate by":
|
||||
return [gr.update(visible=True), gr.update(visible=True), gr.update(visible=False)]
|
||||
else:
|
||||
return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)]
|
||||
resize_method.change(
|
||||
upload_resize_method, resize_method, [width_slider, height_slider, base_resolution]
|
||||
)
|
||||
|
||||
generate_button.click(
|
||||
fn=controller.generate,
|
||||
inputs=[
|
||||
diffusion_transformer_dropdown,
|
||||
base_model_dropdown,
|
||||
lora_model_dropdown,
|
||||
lora_alpha_slider,
|
||||
prompt_textbox,
|
||||
negative_prompt_textbox,
|
||||
sampler_dropdown,
|
||||
sample_step_slider,
|
||||
resize_method,
|
||||
width_slider,
|
||||
height_slider,
|
||||
base_resolution,
|
||||
generation_method,
|
||||
length_slider,
|
||||
cfg_scale_slider,
|
||||
start_image,
|
||||
end_image,
|
||||
validation_video,
|
||||
validation_video_mask,
|
||||
denoise_strength,
|
||||
seed_textbox,
|
||||
ref_image,
|
||||
enable_teacache,
|
||||
teacache_threshold,
|
||||
num_skip_start_steps,
|
||||
teacache_offload,
|
||||
cfg_skip_ratio,
|
||||
enable_riflex,
|
||||
riflex_k,
|
||||
],
|
||||
outputs=[result_image, result_video, infer_progress]
|
||||
)
|
||||
return demo, controller
|
||||
Reference in New Issue
Block a user