Update Wan2.2 UI && Add auto_tile_batch_size args in training && Rewrite Wan2.2 init (#269)
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 = "normal"
|
||||
|
||||
# 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 = "0.0.0.0"
|
||||
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('# --------------------------------------------------------- #')
|
||||
Executable
+192
@@ -0,0 +1,192 @@
|
||||
import base64
|
||||
import json
|
||||
import time
|
||||
import urllib.parse
|
||||
import requests
|
||||
|
||||
|
||||
def post_infer(
|
||||
generation_method,
|
||||
length_slider,
|
||||
url='http://127.0.0.1:7860',
|
||||
POST_TOKEN="",
|
||||
timeout=5,
|
||||
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="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
sampler_dropdown="Flow",
|
||||
sample_step_slider=50,
|
||||
width_slider=672,
|
||||
height_slider=384,
|
||||
cfg_scale_slider=6,
|
||||
seed_textbox=43,
|
||||
enable_teacache = None,
|
||||
teacache_threshold = None,
|
||||
num_skip_start_steps = None,
|
||||
teacache_offload = None,
|
||||
cfg_skip_ratio = None,
|
||||
enable_riflex = None,
|
||||
riflex_k = None,
|
||||
):
|
||||
# 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,
|
||||
|
||||
"enable_teacache": enable_teacache,
|
||||
"teacache_threshold": teacache_threshold,
|
||||
"num_skip_start_steps": num_skip_start_steps,
|
||||
"teacache_offload": teacache_offload,
|
||||
"cfg_skip_ratio": cfg_skip_ratio,
|
||||
"enable_riflex": enable_riflex,
|
||||
"riflex_k": riflex_k,
|
||||
})
|
||||
|
||||
# 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)
|
||||
|
||||
# Extract request ID from POST response headers
|
||||
request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id")
|
||||
|
||||
# Prepare query parameters for GET request
|
||||
query = {
|
||||
'_index_': '0',
|
||||
'_length_': '1',
|
||||
'_timeout_': str(timeout),
|
||||
'_raw_': 'false',
|
||||
'_auto_delete_': 'true',
|
||||
}
|
||||
if request_id:
|
||||
query['requestId'] = request_id
|
||||
|
||||
query_str = urllib.parse.urlencode(query)
|
||||
|
||||
# Polling GET request until status code is not 204
|
||||
status_code = 204
|
||||
while status_code == 204:
|
||||
if query_str:
|
||||
get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout)
|
||||
else:
|
||||
get_r = session.get(f'{url}/sink', timeout=timeout)
|
||||
status_code = get_r.status_code
|
||||
# Decode and return the response content
|
||||
data = get_r.content.decode('utf-8')
|
||||
return data
|
||||
|
||||
if __name__ == '__main__':
|
||||
# initiate time
|
||||
time_start = time.time()
|
||||
|
||||
# EAS队列配置
|
||||
EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx'
|
||||
# Use in EAS Queue
|
||||
TOKEN = 'xxxxxxxx'
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
|
||||
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
|
||||
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
teacache_threshold = 0.10
|
||||
# The number of steps to skip TeaCache at the beginning of the inference process, which can
|
||||
# reduce the impact of TeaCache on generated video quality.
|
||||
num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Skip some cfg steps in inference
|
||||
# Recommended to be set between 0.00 and 0.25
|
||||
cfg_skip_ratio = 0
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# "Video Generation" and "Image Generation"
|
||||
generation_method = "Video Generation"
|
||||
# Video length
|
||||
length_slider = 81
|
||||
# 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 = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
# 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
|
||||
|
||||
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,
|
||||
enable_teacache = enable_teacache,
|
||||
teacache_threshold = teacache_threshold,
|
||||
num_skip_start_steps = num_skip_start_steps,
|
||||
teacache_offload = teacache_offload,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
enable_riflex = enable_riflex,
|
||||
riflex_k = riflex_k,
|
||||
url=EAS_URL,
|
||||
POST_TOKEN=TOKEN
|
||||
)
|
||||
# Get decoded data
|
||||
outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data']))
|
||||
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('# --------------------------------------------------------- #')
|
||||
Executable
+213
@@ -0,0 +1,213 @@
|
||||
import base64
|
||||
import json
|
||||
import time
|
||||
import urllib.parse
|
||||
import requests
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
|
||||
|
||||
def post_infer(
|
||||
generation_method,
|
||||
length_slider,
|
||||
url='http://127.0.0.1:7860',
|
||||
POST_TOKEN="",
|
||||
timeout=5,
|
||||
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="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
sampler_dropdown="Flow",
|
||||
sample_step_slider=50,
|
||||
width_slider=672,
|
||||
height_slider=384,
|
||||
cfg_scale_slider=6,
|
||||
seed_textbox=43,
|
||||
enable_teacache = None,
|
||||
teacache_threshold = None,
|
||||
num_skip_start_steps = None,
|
||||
teacache_offload = None,
|
||||
cfg_skip_ratio = None,
|
||||
enable_riflex = None,
|
||||
riflex_k = None,
|
||||
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,
|
||||
|
||||
"enable_teacache": enable_teacache,
|
||||
"teacache_threshold": teacache_threshold,
|
||||
"num_skip_start_steps": num_skip_start_steps,
|
||||
"teacache_offload": teacache_offload,
|
||||
"cfg_skip_ratio": cfg_skip_ratio,
|
||||
"enable_riflex": enable_riflex,
|
||||
"riflex_k": riflex_k,
|
||||
"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)
|
||||
|
||||
# Extract request ID from POST response headers
|
||||
request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id")
|
||||
|
||||
# Prepare query parameters for GET request
|
||||
query = {
|
||||
'_index_': '0',
|
||||
'_length_': '1',
|
||||
'_timeout_': str(timeout),
|
||||
'_raw_': 'false',
|
||||
'_auto_delete_': 'true',
|
||||
}
|
||||
if request_id:
|
||||
query['requestId'] = request_id
|
||||
|
||||
query_str = urllib.parse.urlencode(query)
|
||||
|
||||
# Polling GET request until status code is not 204
|
||||
status_code = 204
|
||||
while status_code == 204:
|
||||
if query_str:
|
||||
get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout)
|
||||
else:
|
||||
get_r = session.get(f'{url}/sink', timeout=timeout)
|
||||
status_code = get_r.status_code
|
||||
# Decode and return the response content
|
||||
data = get_r.content.decode('utf-8')
|
||||
return data
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# initiate time
|
||||
time_start = time.time()
|
||||
|
||||
# EAS队列配置
|
||||
EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx'
|
||||
# Use in EAS Queue
|
||||
TOKEN = 'xxxxxxxx'
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
|
||||
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
|
||||
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
teacache_threshold = 0.10
|
||||
# The number of steps to skip TeaCache at the beginning of the inference process, which can
|
||||
# reduce the impact of TeaCache on generated video quality.
|
||||
num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Skip some cfg steps in inference
|
||||
# Recommended to be set between 0.00 and 0.25
|
||||
cfg_skip_ratio = 0
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# "Video Generation" and "Image Generation"
|
||||
generation_method = "Video Generation"
|
||||
# Video length
|
||||
length_slider = 81
|
||||
# Used in Lora models
|
||||
lora_model_path = "none"
|
||||
lora_alpha_slider = 0.55
|
||||
# Prompts
|
||||
prompt_textbox = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。"
|
||||
negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
# 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/1.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,
|
||||
enable_teacache = enable_teacache,
|
||||
teacache_threshold = teacache_threshold,
|
||||
num_skip_start_steps = num_skip_start_steps,
|
||||
teacache_offload = teacache_offload,
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
enable_riflex = enable_riflex,
|
||||
riflex_k = riflex_k,
|
||||
url=EAS_URL,
|
||||
POST_TOKEN=TOKEN,
|
||||
start_image=start_image_path # 传递起始图片路径
|
||||
)
|
||||
# Get decoded data
|
||||
outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data']))
|
||||
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('# --------------------------------------------------------- #')
|
||||
@@ -570,6 +570,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1343,7 +1346,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
|
||||
@@ -520,6 +520,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1273,7 +1276,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
control_pixel_values = batch["control_pixel_values"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -582,6 +582,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames."
|
||||
)
|
||||
@@ -1338,7 +1341,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
|
||||
@@ -579,6 +579,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1516,7 +1519,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1537,7 +1540,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -586,6 +586,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames."
|
||||
)
|
||||
@@ -1512,7 +1515,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1533,7 +1536,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -543,6 +543,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1513,7 +1516,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1534,7 +1537,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -459,6 +459,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1524,7 +1527,7 @@ def main():
|
||||
control_camera_values = batch["control_camera_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1))
|
||||
@@ -1551,7 +1554,7 @@ def main():
|
||||
clip_pixel_values = batch["clip_pixel_values"]
|
||||
clip_idx = batch["clip_idx"]
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -477,6 +477,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1528,7 +1531,7 @@ def main():
|
||||
control_camera_values = batch["control_camera_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1))
|
||||
@@ -1555,7 +1558,7 @@ def main():
|
||||
clip_pixel_values = batch["clip_pixel_values"]
|
||||
clip_idx = batch["clip_idx"]
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -555,6 +555,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1512,7 +1515,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1533,7 +1536,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -572,6 +572,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
@@ -1515,7 +1518,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1535,7 +1538,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
mask = torch.tile(mask, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -585,6 +585,9 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames."
|
||||
)
|
||||
@@ -1522,7 +1525,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1542,7 +1545,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.auto_tile_batch_size and args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
mask = torch.tile(mask, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -93,7 +93,9 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller):
|
||||
datas: dict,
|
||||
):
|
||||
base_model_path = datas.get('base_model_path', 'none')
|
||||
base_model_2_path = datas.get('base_model_2_path', 'none')
|
||||
lora_model_path = datas.get('lora_model_path', 'none')
|
||||
lora_model_2_path = datas.get('lora_model_2_path', 'none')
|
||||
lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)
|
||||
prompt_textbox = datas.get('prompt_textbox', None)
|
||||
negative_prompt_textbox = datas.get('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. ')
|
||||
@@ -205,6 +207,8 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller):
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
enable_riflex = enable_riflex,
|
||||
riflex_k = riflex_k,
|
||||
base_model_2_path = base_model_2_path,
|
||||
lora_model_2_path = lora_model_2_path,
|
||||
fps = fps,
|
||||
is_api = True,
|
||||
)
|
||||
|
||||
@@ -99,7 +99,9 @@ if ray is not None:
|
||||
def generate(self, datas):
|
||||
try:
|
||||
base_model_path = datas.get('base_model_path', 'none')
|
||||
base_model_2_path = datas.get('base_model_2_path', 'none')
|
||||
lora_model_path = datas.get('lora_model_path', 'none')
|
||||
lora_model_2_path = datas.get('lora_model_2_path', 'none')
|
||||
lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)
|
||||
prompt_textbox = datas.get('prompt_textbox', None)
|
||||
negative_prompt_textbox = datas.get('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. ')
|
||||
@@ -211,6 +213,8 @@ if ray is not None:
|
||||
cfg_skip_ratio = cfg_skip_ratio,
|
||||
enable_riflex = enable_riflex,
|
||||
riflex_k = riflex_k,
|
||||
base_model_2_path = base_model_2_path,
|
||||
lora_model_2_path = lora_model_2_path,
|
||||
fps = fps,
|
||||
is_api = True,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -81,6 +81,7 @@ class Fun_Controller:
|
||||
self.diffusion_transformer_dropdown = model_name
|
||||
self.scheduler_dict = scheduler_dict
|
||||
self.model_type = model_type
|
||||
self.config_path = os.path.realpath(config_path)
|
||||
if config_path is not None:
|
||||
self.config = OmegaConf.load(config_path)
|
||||
self.ulysses_degree = ulysses_degree
|
||||
@@ -94,21 +95,35 @@ class Fun_Controller:
|
||||
self.diffusion_transformer_list = []
|
||||
self.motion_module_list = []
|
||||
self.personalized_model_list = []
|
||||
self.config_list = []
|
||||
|
||||
# config models
|
||||
self.tokenizer = None
|
||||
self.text_encoder = None
|
||||
self.vae = None
|
||||
self.transformer = None
|
||||
self.transformer_2 = None
|
||||
self.pipeline = None
|
||||
self.base_model_path = "none"
|
||||
self.base_model_2_path = "none"
|
||||
self.lora_model_path = "none"
|
||||
self.lora_model_2_path = "none"
|
||||
|
||||
self.refresh_config()
|
||||
self.refresh_diffusion_transformer()
|
||||
self.refresh_personalized_model()
|
||||
if model_name != None:
|
||||
self.update_diffusion_transformer(model_name)
|
||||
|
||||
def refresh_config(self):
|
||||
config_list = []
|
||||
for root, dirs, files in os.walk(self.config_dir):
|
||||
for file in files:
|
||||
if file.endswith(('.yaml', '.yml')):
|
||||
full_path = os.path.join(root, file)
|
||||
config_list.append(full_path)
|
||||
self.config_list = config_list
|
||||
|
||||
def refresh_diffusion_transformer(self):
|
||||
self.diffusion_transformer_list = sorted(glob(os.path.join(self.diffusion_transformer_dir, "*/")))
|
||||
|
||||
@@ -119,15 +134,27 @@ class Fun_Controller:
|
||||
def update_model_type(self, model_type):
|
||||
self.model_type = model_type
|
||||
|
||||
def update_config(self, config_dropdown):
|
||||
self.config_path = config_dropdown
|
||||
self.config = OmegaConf.load(config_dropdown)
|
||||
print(f"Update config: {config_dropdown}")
|
||||
|
||||
def update_diffusion_transformer(self, diffusion_transformer_dropdown):
|
||||
pass
|
||||
|
||||
def update_base_model(self, base_model_dropdown):
|
||||
self.base_model_path = base_model_dropdown
|
||||
def update_base_model(self, base_model_dropdown, is_checkpoint_2=False):
|
||||
if not is_checkpoint_2:
|
||||
self.base_model_path = base_model_dropdown
|
||||
else:
|
||||
self.base_model_2_path = base_model_dropdown
|
||||
print(f"Update base model: {base_model_dropdown}")
|
||||
if base_model_dropdown == "none":
|
||||
return gr.update()
|
||||
if self.transformer is None:
|
||||
if self.transformer is None and not is_checkpoint_2:
|
||||
gr.Info(f"Please select a pretrained model path.")
|
||||
print(f"Please select a pretrained model path.")
|
||||
return gr.update(value=None)
|
||||
elif self.transformer_2 is None and is_checkpoint_2:
|
||||
gr.Info(f"Please select a pretrained model path.")
|
||||
print(f"Please select a pretrained model path.")
|
||||
return gr.update(value=None)
|
||||
@@ -137,17 +164,23 @@ class Fun_Controller:
|
||||
with safe_open(base_model_dropdown, framework="pt", device="cpu") as f:
|
||||
for key in f.keys():
|
||||
base_model_state_dict[key] = f.get_tensor(key)
|
||||
self.transformer.load_state_dict(base_model_state_dict, strict=False)
|
||||
if not is_checkpoint_2:
|
||||
self.transformer.load_state_dict(base_model_state_dict, strict=False)
|
||||
else:
|
||||
self.transformer_2.load_state_dict(base_model_state_dict, strict=False)
|
||||
print("Update base model done")
|
||||
return gr.update()
|
||||
|
||||
def update_lora_model(self, lora_model_dropdown):
|
||||
def update_lora_model(self, lora_model_dropdown, is_checkpoint_2=False):
|
||||
print(f"Update lora model: {lora_model_dropdown}")
|
||||
if lora_model_dropdown == "none":
|
||||
self.lora_model_path = "none"
|
||||
return gr.update()
|
||||
lora_model_dropdown = os.path.join(self.personalized_model_dir, lora_model_dropdown)
|
||||
self.lora_model_path = lora_model_dropdown
|
||||
if not is_checkpoint_2:
|
||||
self.lora_model_path = lora_model_dropdown
|
||||
else:
|
||||
self.lora_model_2_path = lora_model_dropdown
|
||||
return gr.update()
|
||||
|
||||
def clear_cache(self,):
|
||||
|
||||
+40
-2
@@ -79,7 +79,7 @@ def create_fake_model_checkpoints(model_name, visible):
|
||||
)
|
||||
return diffusion_transformer_dropdown
|
||||
|
||||
def create_finetune_models_checkpoints(controller, visible):
|
||||
def create_finetune_models_checkpoints(controller, visible, add_checkpoint_2=False):
|
||||
with gr.Row(visible=visible):
|
||||
base_model_dropdown = gr.Dropdown(
|
||||
label="Select base Dreambooth model (选择基模型[非必需])",
|
||||
@@ -87,6 +87,13 @@ def create_finetune_models_checkpoints(controller, visible):
|
||||
value="none",
|
||||
interactive=True,
|
||||
)
|
||||
if add_checkpoint_2:
|
||||
base_model_2_dropdown = gr.Dropdown(
|
||||
label="Select base Dreambooth model (选择第二个基模型[非必需])",
|
||||
choices=["none"] + controller.personalized_model_list,
|
||||
value="none",
|
||||
interactive=True,
|
||||
)
|
||||
|
||||
lora_model_dropdown = gr.Dropdown(
|
||||
label="Select LoRA model (选择LoRA模型[非必需])",
|
||||
@@ -94,6 +101,13 @@ def create_finetune_models_checkpoints(controller, visible):
|
||||
value="none",
|
||||
interactive=True,
|
||||
)
|
||||
if add_checkpoint_2:
|
||||
lora_model_2_dropdown = gr.Dropdown(
|
||||
label="Select LoRA model (选择LoRA模型[非必需])",
|
||||
choices=["none"] + controller.personalized_model_list,
|
||||
value="none",
|
||||
interactive=True,
|
||||
)
|
||||
|
||||
lora_alpha_slider = gr.Slider(label="LoRA alpha (LoRA权重)", value=0.55, minimum=0, maximum=2, interactive=True)
|
||||
|
||||
@@ -106,7 +120,11 @@ def create_finetune_models_checkpoints(controller, visible):
|
||||
]
|
||||
personalized_refresh_button.click(fn=update_personalized_model, inputs=[], outputs=[base_model_dropdown, lora_model_dropdown])
|
||||
|
||||
return base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button
|
||||
if not add_checkpoint_2:
|
||||
return base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button
|
||||
else:
|
||||
return [base_model_dropdown, base_model_2_dropdown], [lora_model_dropdown, lora_model_2_dropdown], \
|
||||
lora_alpha_slider, personalized_refresh_button
|
||||
|
||||
def create_fake_finetune_models_checkpoints(visible):
|
||||
with gr.Row():
|
||||
@@ -318,3 +336,23 @@ def create_ui_outputs():
|
||||
interactive=False
|
||||
)
|
||||
return result_image, result_video, infer_progress
|
||||
|
||||
def create_config(controller):
|
||||
gr.Markdown(
|
||||
"""
|
||||
### Config Path (配置文件路径)
|
||||
"""
|
||||
)
|
||||
with gr.Row():
|
||||
config_dropdown = gr.Dropdown(
|
||||
label="Config Path (配置文件路径)",
|
||||
choices=controller.config_list,
|
||||
value=controller.config_path,
|
||||
interactive=True,
|
||||
)
|
||||
config_refresh_button = gr.Button(value="\U0001F503", elem_classes="toolbutton")
|
||||
def refresh_config():
|
||||
controller.refresh_config()
|
||||
return gr.update(choices=controller.config_list)
|
||||
config_refresh_button.click(fn=refresh_config, inputs=[], outputs=[config_dropdown])
|
||||
return config_dropdown, config_refresh_button
|
||||
@@ -0,0 +1,766 @@
|
||||
"""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, create_config)
|
||||
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,
|
||||
base_model_2_dropdown=None,
|
||||
lora_model_2_dropdown=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.base_model_2_path != base_model_2_dropdown:
|
||||
self.update_lora_model(base_model_2_dropdown, is_checkpoint_2=True)
|
||||
|
||||
if self.lora_model_path != lora_model_dropdown:
|
||||
self.update_lora_model(lora_model_dropdown)
|
||||
if self.lora_model_2_path != lora_model_2_dropdown:
|
||||
self.update_lora_model(lora_model_2_dropdown, is_checkpoint_2=True)
|
||||
|
||||
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)
|
||||
self.pipeline = merge_lora(self.pipeline, self.lora_model_2_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2")
|
||||
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_2_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_2_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"):
|
||||
config_dropdown, config_refresh_button = create_config(controller)
|
||||
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, add_checkpoint_2=True)
|
||||
base_model_dropdown, base_model_2_dropdown = base_model_dropdown
|
||||
lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown
|
||||
|
||||
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()
|
||||
|
||||
config_dropdown.change(
|
||||
fn=controller.update_config,
|
||||
inputs=[config_dropdown],
|
||||
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,
|
||||
base_model_2_dropdown,
|
||||
lora_model_2_dropdown
|
||||
],
|
||||
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, add_checkpoint_2=True)
|
||||
base_model_dropdown, base_model_2_dropdown = base_model_dropdown
|
||||
lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown
|
||||
|
||||
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,
|
||||
base_model_2_dropdown,
|
||||
lora_model_2_dropdown
|
||||
],
|
||||
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, add_checkpoint_2=True)
|
||||
base_model_dropdown, base_model_2_dropdown = base_model_dropdown
|
||||
lora_model_dropdown, lora_model_2_dropdown = lora_model_dropdown
|
||||
|
||||
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,
|
||||
base_model_2_dropdown,
|
||||
lora_model_2_dropdown
|
||||
],
|
||||
outputs=[result_image, result_video, infer_progress]
|
||||
)
|
||||
return demo, controller
|
||||
Reference in New Issue
Block a user