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('# --------------------------------------------------------- #')
|
||||
Reference in New Issue
Block a user