Update Wan2.2 UI && Add auto_tile_batch_size args in training && Rewrite Wan2.2 init (#269)

This commit is contained in:
Bubbliiiing
2025-07-30 19:54:27 +08:00
committed by GitHub
parent 34ef44eb52
commit 4dc94d2b95
22 changed files with 1678 additions and 91 deletions
+79
View File
@@ -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)
+91
View File
@@ -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()
+169
View File
@@ -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('# --------------------------------------------------------- #')
+192
View File
@@ -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('# --------------------------------------------------------- #')
+213
View File
@@ -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('# --------------------------------------------------------- #')
+4 -1
View File
@@ -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:
+4 -1
View File
@@ -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))
+4 -1
View File
@@ -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:
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+4
View File
@@ -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,
)
+4
View File
@@ -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,
)
+29 -64
View File
@@ -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
+39 -6
View File
@@ -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
View File
@@ -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
+766
View File
@@ -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