diff --git a/comfyui/wan2_1/nodes.py b/comfyui/wan2_1/nodes.py index 52bb724..c2d6c83 100755 --- a/comfyui/wan2_1/nodes.py +++ b/comfyui/wan2_1/nodes.py @@ -352,6 +352,7 @@ class WanT2VSampler: video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 if riflex_k > 0: + latent_frames = (video_length - 1) // self.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) # Apply lora @@ -510,6 +511,7 @@ class WanI2VSampler: input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width)) if riflex_k > 0: + latent_frames = (video_length - 1) // self.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) # Apply lora diff --git a/comfyui/wan2_1_fun/nodes.py b/comfyui/wan2_1_fun/nodes.py index 6c04eb5..c9fccce 100755 --- a/comfyui/wan2_1_fun/nodes.py +++ b/comfyui/wan2_1_fun/nodes.py @@ -357,6 +357,7 @@ class WanFunT2VSampler: video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 if riflex_k > 0: + latent_frames = (video_length - 1) // self.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) # Apply lora @@ -532,6 +533,7 @@ class WanFunInpaintSampler: video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 if riflex_k > 0: + latent_frames = (video_length - 1) // self.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width)) @@ -718,6 +720,7 @@ class WanFunV2VSampler: video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 if riflex_k > 0: + latent_frames = (video_length - 1) // self.vae.config.temporal_compression_ratio + 1 pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) if model_type == "Inpaint": diff --git a/examples/cogvideox_fun/launch_api.py b/examples/cogvideox_fun/launch_api.py index aef0191..b7b410f 100755 --- a/examples/cogvideox_fun/launch_api.py +++ b/examples/cogvideox_fun/launch_api.py @@ -28,6 +28,7 @@ def main(): parser.add_argument('--server_port', type=int, default=7860, help='Server Port') parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP", 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 @@ -40,7 +41,7 @@ def main(): world_size=args.world_size, Controller=CogVideoXFunController, GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=None, ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree, enable_teacache=False, teacache_threshold=0.1, num_skip_start_steps=5, - teacache_offload=False, weight_dtype=weight_dtype, + teacache_offload=False, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample, ) def gr_launch(): diff --git a/examples/wan2.1/app.py b/examples/wan2.1/app.py index a613edb..93bd13f 100755 --- a/examples/wan2.1/app.py +++ b/examples/wan2.1/app.py @@ -67,11 +67,11 @@ if __name__ == "__main__": model_type = "Inpaint" if ui_mode == "host": - demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype) + demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, 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, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype) + demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype) def gr_launch(): # launch gradio diff --git a/examples/wan2.1/launch_api.py b/examples/wan2.1/launch_api.py index 3c682e0..c451e4e 100755 --- a/examples/wan2.1/launch_api.py +++ b/examples/wan2.1/launch_api.py @@ -33,6 +33,7 @@ def main(): parser.add_argument('--config_path', type=str, default="config/wan2.1/wan_civitai.yaml", help='Path to config file') parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/Wan2.1-T2V-1.3B", 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 @@ -45,7 +46,7 @@ def main(): world_size=args.world_size, Controller=Wan_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, enable_teacache=args.enable_teacache, teacache_threshold=args.teacache_threshold, num_skip_start_steps=args.num_skip_start_steps, - teacache_offload=args.teacache_offload, weight_dtype=weight_dtype, + teacache_offload=args.teacache_offload, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample, ) def gr_launch(): diff --git a/examples/wan2.1/post_infer_queue.py b/examples/wan2.1/post_infer_queue.py index 374c92b..c046738 100755 --- a/examples/wan2.1/post_infer_queue.py +++ b/examples/wan2.1/post_infer_queue.py @@ -15,7 +15,7 @@ def post_infer( 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.", + negative_prompt_textbox="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", sampler_dropdown="Flow", sample_step_slider=50, width_slider=672, @@ -87,13 +87,13 @@ if __name__ == '__main__': # "Video Generation" and "Image Generation" generation_method = "Video Generation" # Video length - length_slider = 49 + 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 = "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." + negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # Sampler name sampler_dropdown = "Flow" # Sampler steps diff --git a/examples/wan2.1/post_infer_queue_i2v.py b/examples/wan2.1/post_infer_queue_i2v.py new file mode 100755 index 0000000..14145cb --- /dev/null +++ b/examples/wan2.1/post_infer_queue_i2v.py @@ -0,0 +1,163 @@ +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, + start_image=None +): + if start_image: + try: + image = Image.open(start_image) + # 将图片转换为 Base64 编码 + buffered = BytesIO() + image.save(buffered, format=image.format) + 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 + 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' + + # "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, + 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) % 60 + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') diff --git a/examples/wan2.1_fun/launch_api.py b/examples/wan2.1_fun/launch_api.py index db72708..f9d737a 100755 --- a/examples/wan2.1_fun/launch_api.py +++ b/examples/wan2.1_fun/launch_api.py @@ -33,6 +33,7 @@ def main(): parser.add_argument('--config_path', type=str, default="config/wan2.1/wan_civitai.yaml", help='Path to config file') parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/Wan2.1-Fun-1.3B-InP", 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 @@ -45,7 +46,7 @@ def main(): world_size=args.world_size, Controller=Wan_Fun_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, enable_teacache=args.enable_teacache, teacache_threshold=args.teacache_threshold, num_skip_start_steps=args.num_skip_start_steps, - teacache_offload=args.teacache_offload, weight_dtype=weight_dtype, + teacache_offload=args.teacache_offload, weight_dtype=weight_dtype, savedir_sample=args.savedir_sample, ) def gr_launch(): diff --git a/examples/wan2.1_fun/post_infer_queue.py b/examples/wan2.1_fun/post_infer_queue.py index 374c92b..c046738 100755 --- a/examples/wan2.1_fun/post_infer_queue.py +++ b/examples/wan2.1_fun/post_infer_queue.py @@ -15,7 +15,7 @@ def post_infer( 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.", + negative_prompt_textbox="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", sampler_dropdown="Flow", sample_step_slider=50, width_slider=672, @@ -87,13 +87,13 @@ if __name__ == '__main__': # "Video Generation" and "Image Generation" generation_method = "Video Generation" # Video length - length_slider = 49 + 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 = "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." + negative_prompt_textbox = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # Sampler name sampler_dropdown = "Flow" # Sampler steps diff --git a/examples/wan2.1_fun/post_infer_queue_i2v.py b/examples/wan2.1_fun/post_infer_queue_i2v.py new file mode 100755 index 0000000..14145cb --- /dev/null +++ b/examples/wan2.1_fun/post_infer_queue_i2v.py @@ -0,0 +1,163 @@ +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, + start_image=None +): + if start_image: + try: + image = Image.open(start_image) + # 将图片转换为 Base64 编码 + buffered = BytesIO() + image.save(buffered, format=image.format) + 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 + 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' + + # "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, + 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) % 60 + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') diff --git a/examples/wan2.1_fun/post_infer_queue_v2v_control.py b/examples/wan2.1_fun/post_infer_queue_v2v_control.py new file mode 100755 index 0000000..f3bff0a --- /dev/null +++ b/examples/wan2.1_fun/post_infer_queue_v2v_control.py @@ -0,0 +1,160 @@ +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, + control_video=None +): + if control_video: + try: + if not control_video.startswith("http"): + with open(control_video, "rb") as file: + video_data = file.read() + + control_video = base64.b64encode(video_data).decode('utf-8') + except Exception as e: + print(f"Error processing control_video: {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, + "control_video": control_video + }) + + # Initialize session and set headers + session = requests.session() + session.headers.update({"Authorization": POST_TOKEN}) + + # Send POST request + 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' + + # "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 = 480 + height_slider = 832 + # cfg scale + cfg_scale_slider = 6 + seed_textbox = 43 + + # 控制视频路径(可以是本地路径或 URL) + control_video_path = "asset/000000.mp4" # 替换为实际的视频路径 + + 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=EAS_URL, + POST_TOKEN=TOKEN, + control_video=control_video_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) % 60 + print('# --------------------------------------------------------- #') + print(f'# Total expenditure: {time_sum}s') + print('# --------------------------------------------------------- #') diff --git a/scripts/zero_to_bf16.py b/scripts/zero_to_bf16.py old mode 100644 new mode 100755 index 0bc3543..47623dc --- a/scripts/zero_to_bf16.py +++ b/scripts/zero_to_bf16.py @@ -16,24 +16,31 @@ # python zero_to_bf16.py . output_dir/ --safe_serialization import argparse -import torch +import gc import glob +import json import math import os +import queue import re -import gc -import json -import numpy as np -from tqdm import tqdm from collections import OrderedDict from dataclasses import dataclass +from threading import Thread +import numpy as np +import torch +from deepspeed.checkpoint.constants import (BUFFER_NAMES, DS_VERSION, + FP32_FLAT_GROUPS, + FROZEN_PARAM_FRAGMENTS, + FROZEN_PARAM_SHAPES, + OPTIMIZER_STATE_DICT, PARAM_SHAPES, + PARTITION_COUNT, + SINGLE_PARTITION_OF_FP32_GROUPS, + ZERO_STAGE) # while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with # DeepSpeed data structures it has to be available in the current python environment. from deepspeed.utils import logger -from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, - FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, - FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) +from tqdm import tqdm @dataclass @@ -516,17 +523,35 @@ def to_torch_tensor(state_dict, return_empty_tensor=False): """ torch_state_dict = {} converted_tensors = {} - for name, tensor in state_dict.items(): - tensor_id = id(tensor) - if tensor_id in converted_tensors: # shared tensors - shared_tensor = torch_state_dict[converted_tensors[tensor_id]] - torch_state_dict[name] = shared_tensor - else: - converted_tensors[tensor_id] = name - if return_empty_tensor: - torch_state_dict[name] = torch.empty(tensor.shape, dtype=tensor.dtype) + def convert_tensor(qin): + while True: + name, tensor = qin.get() + if name is None: + return + tensor_id = id(tensor) + if tensor_id in converted_tensors: + shared_tensor = torch_state_dict[converted_tensors[tensor_id]] + torch_state_dict[name] = shared_tensor.to(torch.bfloat16) else: - torch_state_dict[name] = tensor.contiguous() + converted_tensors[tensor_id] = name + if return_empty_tensor: + torch_state_dict[name] = torch.empty(tensor.shape, dtype=tensor.dtype).to(torch.bfloat16) + else: + torch_state_dict[name] = tensor.contiguous().to(torch.bfloat16) + + num_threads = 32 + qin = queue.Queue(num_threads) + threads = [Thread(target=convert_tensor, args=(qin, )) for _ in range(num_threads)] + [_.start() for _ in threads] + cnt = 0 + for name, tensor in state_dict.items(): + cnt += 1 + qin.put([name, tensor]) + if cnt % 1000 == 0: + print(f'{cnt} / {len(state_dict)}') + for _ in range(num_threads): + qin.put([None, None]) + [_.join() for _ in threads] return torch_state_dict @@ -655,7 +680,11 @@ def convert_zero_checkpoint_to_bf16_state_dict(checkpoint_dir, for shard_file, tensors in tqdm(filename_to_tensors, desc="Saving checkpoint shards"): shard_state_dict = {tensor_name: state_dict[tensor_name] for tensor_name in tensors} shard_state_dict = to_torch_tensor(shard_state_dict) - shard_state_dict = {tensor_name: shard_state_dict[tensor_name].to(torch.bfloat16) for tensor_name in shard_state_dict} + # to bf16 + shard_state_dict = { + tensor_name: shard_state_dict[tensor_name].to(torch.bfloat16) for tensor_name in shard_state_dict + } + print('save shard_state_dict') output_path = os.path.join(output_dir, shard_file) if safe_serialization: save_file(shard_state_dict, output_path, metadata={"format": "pt"}) diff --git a/videox_fun/api/api.py b/videox_fun/api/api.py index e1f1d0d..e525fda 100755 --- a/videox_fun/api/api.py +++ b/videox_fun/api/api.py @@ -1,16 +1,18 @@ -import io -import gc import base64 -import torch -import gradio as gr -import tempfile +import gc import hashlib +import io import os - -from fastapi import FastAPI +import tempfile from io import BytesIO + +import gradio as gr +import requests +import torch +from fastapi import FastAPI from PIL import Image + # Function to encode a file to Base64 def encode_file_to_base64(file_path): with open(file_path, "rb") as file: @@ -54,6 +56,15 @@ def update_diffusion_transformer_api(_: gr.Blocks, app: FastAPI, controller): return {"message": comment} +def download_from_url(url, timeout=10): + try: + response = requests.get(url, timeout=timeout) + response.raise_for_status() # 检查请求是否成功 + return response.content + except requests.exceptions.RequestException as e: + print(f"Error downloading from {url}: {e}") + return None + def save_base64_video(base64_string): video_data = base64.b64decode(base64_string) @@ -82,6 +93,18 @@ def save_base64_image(base64_string): return file_path +def save_url_video(url): + video_data = download_from_url(url) + if video_data: + return save_base64_video(base64.b64encode(video_data)) + return None + +def save_url_image(url): + image_data = download_from_url(url) + if image_data: + return save_base64_image(base64.b64encode(image_data)) + return None + def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): @app.post("/videox_fun/infer_forward") def _infer_forward_api( @@ -115,21 +138,38 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): generation_method = "Image Generation" if is_image else generation_method if start_image is not None: - start_image = base64.b64decode(start_image) - start_image = [Image.open(BytesIO(start_image))] - + if start_image.startswith('http'): + start_image = save_url_image(start_image) + start_image = [Image.open(start_image)] + else: + start_image = base64.b64decode(start_image) + start_image = [Image.open(BytesIO(start_image))] + if end_image is not None: - end_image = base64.b64decode(end_image) - end_image = [Image.open(BytesIO(end_image))] + if end_image.startswith('http'): + end_image = save_url_image(end_image) + end_image = [Image.open(end_image)] + else: + end_image = base64.b64decode(end_image) + end_image = [Image.open(BytesIO(end_image))] if validation_video is not None: - validation_video = save_base64_video(validation_video) + if validation_video.startswith('http'): + validation_video = save_url_video(validation_video) + else: + validation_video = save_base64_video(validation_video) if validation_video_mask is not None: - validation_video_mask = save_base64_image(validation_video_mask) + if validation_video_mask.startswith('http'): + validation_video_mask = save_url_image(validation_video_mask) + else: + validation_video_mask = save_base64_image(validation_video_mask) if control_video is not None: - control_video = save_base64_video(control_video) + if control_video.startswith('http'): + control_video = save_url_video(control_video) + else: + control_video = save_base64_video(control_video) try: save_sample_path, comment = controller.generate( diff --git a/videox_fun/api/api_multi_nodes.py b/videox_fun/api/api_multi_nodes.py index 5dbcd52..66e8d59 100755 --- a/videox_fun/api/api_multi_nodes.py +++ b/videox_fun/api/api_multi_nodes.py @@ -9,7 +9,8 @@ import torch from fastapi import FastAPI, HTTPException from PIL import Image -from .api import encode_file_to_base64, save_base64_image, save_base64_video +from .api import (encode_file_to_base64, save_base64_image, save_base64_video, + save_url_image, save_url_video) try: import ray @@ -26,6 +27,7 @@ if ray is not None: config_path=None, ulysses_degree=1, ring_degree=1, enable_teacache=None, teacache_threshold=None, num_skip_start_steps=None, teacache_offload=None, weight_dtype=None, + savedir_sample=None, ): # Set PyTorch distributed environment variables os.environ["RANK"] = str(rank) @@ -37,7 +39,7 @@ if ray is not None: self.controller = Controller( GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, - teacache_offload=teacache_offload, weight_dtype=weight_dtype, + teacache_offload=teacache_offload, weight_dtype=weight_dtype, savedir_sample=savedir_sample, ) def generate(self, datas): @@ -70,21 +72,38 @@ if ray is not None: generation_method = "Image Generation" if is_image else generation_method if start_image is not None: - start_image = base64.b64decode(start_image) - start_image = [Image.open(BytesIO(start_image))] - - if end_image is not None: - end_image = base64.b64decode(end_image) - end_image = [Image.open(BytesIO(end_image))] + if start_image.startswith('http'): + start_image = save_url_image(start_image) + start_image = [Image.open(start_image)] + else: + start_image = base64.b64decode(start_image) + start_image = [Image.open(BytesIO(start_image))] + if end_image is not None: + if end_image.startswith('http'): + end_image = save_url_image(end_image) + end_image = [Image.open(end_image)] + else: + end_image = base64.b64decode(end_image) + end_image = [Image.open(BytesIO(end_image))] + if validation_video is not None: - validation_video = save_base64_video(validation_video) + if validation_video.startswith('http'): + validation_video = save_url_video(validation_video) + else: + validation_video = save_base64_video(validation_video) if validation_video_mask is not None: - validation_video_mask = save_base64_image(validation_video_mask) + if validation_video_mask.startswith('http'): + validation_video_mask = save_url_image(validation_video_mask) + else: + validation_video_mask = save_base64_image(validation_video_mask) if control_video is not None: - control_video = save_base64_video(control_video) + if control_video.startswith('http'): + control_video = save_url_video(control_video) + else: + control_video = save_base64_video(control_video) try: save_sample_path, comment = self.controller.generate( @@ -150,7 +169,8 @@ if ray is not None: teacache_threshold, num_skip_start_steps, teacache_offload, - weight_dtype + weight_dtype, + savedir_sample ): # Ensure Ray is initialized if not ray.is_initialized(): @@ -162,7 +182,7 @@ if ray is not None: rank, world_size, Controller, GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, - teacache_offload=teacache_offload, weight_dtype=weight_dtype, + teacache_offload=teacache_offload, weight_dtype=weight_dtype, savedir_sample=savedir_sample, ) for rank in range(num_workers) ] diff --git a/videox_fun/data/bucket_sampler.py b/videox_fun/data/bucket_sampler.py old mode 100644 new mode 100755 index 2c5fded..24b4160 --- a/videox_fun/data/bucket_sampler.py +++ b/videox_fun/data/bucket_sampler.py @@ -254,7 +254,7 @@ class AspectRatioBatchSampler(BatchSampler): width = int(width) ratio = height / width # self.dataset[idx] except Exception as e: - print(e) + print(e, self.dataset[idx], "This item is error, please check it.") continue # find the closest aspect ratio closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio)) @@ -330,7 +330,7 @@ class AspectRatioBatchImageVideoSampler(BatchSampler): width = int(width) ratio = height / width # self.dataset[idx] except Exception as e: - print(e) + print(e, self.dataset[idx], "This item is error, please check it.") continue # find the closest aspect ratio closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio)) @@ -365,7 +365,7 @@ class AspectRatioBatchImageVideoSampler(BatchSampler): width = int(width) ratio = height / width # self.dataset[idx] except Exception as e: - print(e) + print(e, self.dataset[idx], "This item is error, please check it.") continue # find the closest aspect ratio closest_ratio = min(self.aspect_ratios.keys(), key=lambda r: abs(float(r) - ratio)) diff --git a/videox_fun/models/wan_text_encoder.py b/videox_fun/models/wan_text_encoder.py old mode 100644 new mode 100755 index 34a0323..49fc936 --- a/videox_fun/models/wan_text_encoder.py +++ b/videox_fun/models/wan_text_encoder.py @@ -362,4 +362,15 @@ class WanT5EncoderModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): except Exception as e: print( f"The low_cpu_mem_usage mode is not work because {e}. Use low_cpu_mem_usage=False instead." - ) \ No newline at end of file + ) + + model = cls(**filter_kwargs(cls, additional_kwargs)) + if pretrained_model_path.endswith(".safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(pretrained_model_path) + else: + state_dict = torch.load(pretrained_model_path, map_location="cpu") + m, u = model.load_state_dict(state_dict, strict=False) + print(f"### missing keys: {len(m)}; \n### unexpected keys: {len(u)};") + print(m, u) + return model \ No newline at end of file diff --git a/videox_fun/ui/cogvideox_fun_ui.py b/videox_fun/ui/cogvideox_fun_ui.py index aa06726..9f87d39 100755 --- a/videox_fun/ui/cogvideox_fun_ui.py +++ b/videox_fun/ui/cogvideox_fun_ui.py @@ -306,11 +306,12 @@ class CogVideoXFunController(Fun_Controller): CogVideoXFunController_Host = CogVideoXFunController CogVideoXFunController_Client = Fun_Controller_Client -def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtype): +def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None): controller = CogVideoXFunController( GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint", ulysses_degree=ulysses_degree, ring_degree=ring_degree, config_path=None, enable_teacache=None, teacache_threshold=None, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -436,11 +437,12 @@ def ui(GPU_memory_mode, scheduler_dict, ulysses_degree, ring_degree, weight_dtyp ) return demo, controller -def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, ulysses_degree, ring_degree, weight_dtype): +def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, ulysses_degree, ring_degree, weight_dtype, savedir_sample=None): controller = CogVideoXFunController_Host( GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, ulysses_degree=ulysses_degree, ring_degree=ring_degree, config_path=None, enable_teacache=None, teacache_threshold=None, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -556,8 +558,8 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, ulysses_deg ) return demo, controller -def ui_client(scheduler_dict, model_name): - controller = CogVideoXFunController_Client(scheduler_dict) +def ui_client(scheduler_dict, model_name, savedir_sample=None): + controller = CogVideoXFunController_Client(scheduler_dict, savedir_sample) with gr.Blocks(css=css) as demo: gr.Markdown( diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py index d798b47..2e87710 100755 --- a/videox_fun/ui/controller.py +++ b/videox_fun/ui/controller.py @@ -58,7 +58,7 @@ class Fun_Controller: config_path=None, ulysses_degree=1, ring_degree=1, enable_teacache=None, teacache_threshold=None, num_skip_start_steps=None, teacache_offload=None, - enable_riflex=None, riflex_k=None, weight_dtype=None, + enable_riflex=None, riflex_k=None, weight_dtype=None, savedir_sample=None, ): # config dirs self.basedir = os.getcwd() @@ -66,9 +66,11 @@ class Fun_Controller: self.diffusion_transformer_dir = os.path.join(self.basedir, "models", "Diffusion_Transformer") self.motion_module_dir = os.path.join(self.basedir, "models", "Motion_Module") self.personalized_model_dir = os.path.join(self.basedir, "models", "Personalized_Model") - self.savedir = os.path.join(self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S")) - self.savedir_sample = os.path.join(self.savedir, "sample") - os.makedirs(self.savedir, exist_ok=True) + if savedir_sample is None: + self.savedir_sample = os.path.join(self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S")) + else: + self.savedir_sample = savedir_sample + os.makedirs(self.savedir_sample, exist_ok=True) self.GPU_memory_mode = GPU_memory_mode self.model_name = model_name @@ -344,10 +346,13 @@ def post_to_host( class Fun_Controller_Client: - def __init__(self, scheduler_dict): + def __init__(self, scheduler_dict, savedir_sample): self.basedir = os.getcwd() - self.savedir = os.path.join(self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S")) - self.savedir_sample = os.path.join(self.savedir, "sample") + if savedir_sample is None: + self.savedir_sample = os.path.join(self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S")) + else: + self.savedir_sample = savedir_sample + os.makedirs(self.savedir_sample, exist_ok=True) self.scheduler_dict = scheduler_dict diff --git a/videox_fun/ui/wan_fun_ui.py b/videox_fun/ui/wan_fun_ui.py index 61acb0e..bc67205 100755 --- a/videox_fun/ui/wan_fun_ui.py +++ b/videox_fun/ui/wan_fun_ui.py @@ -15,7 +15,7 @@ from ..data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio from ..models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, WanT5EncoderModel, WanTransformer3DModel) from ..models.cache_utils import get_teacache_coefficients -from ..pipeline import WanFunInpaintPipeline, WanFunPipeline +from ..pipeline import WanFunInpaintPipeline, WanFunPipeline, WanFunControlPipeline from ..utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper, replace_parameters_by_name) @@ -104,7 +104,14 @@ class Wan_Fun_Controller(Fun_Controller): scheduler=self.scheduler, ) else: - raise ValueError("Not support now") + self.pipeline = WanFunControlPipeline( + vae=self.vae, + tokenizer=self.tokenizer, + text_encoder=self.text_encoder, + transformer=self.transformer, + scheduler=self.scheduler, + clip_image_encoder=self.clip_image_encoder, + ) if self.ulysses_degree > 1 or self.ring_degree > 1: self.transformer.enable_multi_gpus_inference() @@ -187,7 +194,8 @@ class Wan_Fun_Controller(Fun_Controller): generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) if self.enable_riflex: - self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = length_slider if not is_image else 1) + latent_frames = (int(length_slider) - 1) // self.vae.config.temporal_compression_ratio + 1 + self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = latent_frames if not is_image else 1) try: if self.model_type == "Inpaint": @@ -275,13 +283,14 @@ class Wan_Fun_Controller(Fun_Controller): Wan_Fun_Controller_Host = Wan_Fun_Controller Wan_Fun_Controller_Client = Fun_Controller_Client -def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype): +def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype, savedir_sample=None): controller = Wan_Fun_Controller( GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint", config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -401,13 +410,14 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree ) return demo, controller -def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype): +def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype, savedir_sample=None): controller = Wan_Fun_Controller_Host( GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -517,8 +527,8 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path ) return demo, controller -def ui_client(scheduler_dict, model_name): - controller = Wan_Fun_Controller_Client(scheduler_dict) +def ui_client(scheduler_dict, model_name, savedir_sample=None): + controller = Wan_Fun_Controller_Client(scheduler_dict, savedir_sample) with gr.Blocks(css=css) as demo: gr.Markdown( diff --git a/videox_fun/ui/wan_ui.py b/videox_fun/ui/wan_ui.py index c59750a..b99b118 100755 --- a/videox_fun/ui/wan_ui.py +++ b/videox_fun/ui/wan_ui.py @@ -187,7 +187,8 @@ class Wan_Controller(Fun_Controller): generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox)) if self.enable_riflex: - self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = length_slider if not is_image else 1) + latent_frames = (int(length_slider) - 1) // self.vae.config.temporal_compression_ratio + 1 + self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = latent_frames if not is_image else 1) try: if self.model_type == "Inpaint": @@ -275,13 +276,14 @@ class Wan_Controller(Fun_Controller): Wan_Controller_Host = Wan_Controller Wan_Controller_Client = Fun_Controller_Client -def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype): +def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype, savedir_sample=None): controller = Wan_Controller( GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint", config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -397,13 +399,14 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree ) return demo, controller -def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype): +def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype, savedir_sample=None): controller = Wan_Controller_Host( GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree, enable_teacache=enable_teacache, teacache_threshold=teacache_threshold, num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype, + savedir_sample=savedir_sample, ) with gr.Blocks(css=css) as demo: @@ -509,8 +512,8 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path ) return demo, controller -def ui_client(scheduler_dict, model_name): - controller = Wan_Controller_Client(scheduler_dict) +def ui_client(scheduler_dict, model_name, savedir_sample=None): + controller = Wan_Controller_Client(scheduler_dict, savedir_sample) with gr.Blocks(css=css) as demo: gr.Markdown(