diff --git a/videox_fun/api/api.py b/videox_fun/api/api.py index f520d82..253a1d8 100755 --- a/videox_fun/api/api.py +++ b/videox_fun/api/api.py @@ -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, ) diff --git a/videox_fun/api/api_multi_nodes.py b/videox_fun/api/api_multi_nodes.py index 9562932..aaf9d15 100755 --- a/videox_fun/api/api_multi_nodes.py +++ b/videox_fun/api/api_multi_nodes.py @@ -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, ) diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py index 28c2504..d5a16ae 100755 --- a/videox_fun/ui/controller.py +++ b/videox_fun/ui/controller.py @@ -100,9 +100,12 @@ class Fun_Controller: 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_diffusion_transformer() self.refresh_personalized_model() @@ -122,12 +125,19 @@ class Fun_Controller: 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 +147,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,): diff --git a/videox_fun/ui/ui.py b/videox_fun/ui/ui.py index e357716..a24ea2d 100755 --- a/videox_fun/ui/ui.py +++ b/videox_fun/ui/ui.py @@ -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(): diff --git a/videox_fun/ui/wan2_2_ui.py b/videox_fun/ui/wan2_2_ui.py index 712387e..0de7780 100644 --- a/videox_fun/ui/wan2_2_ui.py +++ b/videox_fun/ui/wan2_2_ui.py @@ -187,6 +187,8 @@ class Wan2_2_Controller(Fun_Controller): 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, ): @@ -203,9 +205,13 @@ class Wan2_2_Controller(Fun_Controller): 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 @@ -223,6 +229,7 @@ class Wan2_2_Controller(Fun_Controller): 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 @@ -231,9 +238,7 @@ class Wan2_2_Controller(Fun_Controller): 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 - ) + self.pipeline.transformer_2.share_teacache(self.pipeline.transformer) else: print(f"Disable TeaCache.") self.pipeline.transformer.disable_teacache() @@ -329,7 +334,7 @@ class Wan2_2_Controller(Fun_Controller): print(f"Error. error information is {str(e)}") if self.lora_model_path != "none": self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) - self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2") + 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: @@ -340,7 +345,7 @@ class Wan2_2_Controller(Fun_Controller): if self.lora_model_path != "none": print(f"Unmerge Lora.") self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider) - self.pipeline = unmerge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider, sub_transformer_name="transformer_2") + 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.") @@ -387,7 +392,9 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, 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) + 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 = \ @@ -498,6 +505,8 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype, cfg_skip_ratio, enable_riflex, riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown ], outputs=[result_image, result_video, infer_progress] ) @@ -519,7 +528,10 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path 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) + 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 = \ @@ -622,6 +634,8 @@ def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path cfg_skip_ratio, enable_riflex, riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown ], outputs=[result_image, result_video, infer_progress] ) @@ -638,7 +652,10 @@ def ui_client(scheduler_dict, model_name, savedir_sample=None): ) 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) + 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 = \ @@ -734,6 +751,8 @@ def ui_client(scheduler_dict, model_name, savedir_sample=None): cfg_skip_ratio, enable_riflex, riflex_k, + base_model_2_dropdown, + lora_model_2_dropdown ], outputs=[result_image, result_video, infer_progress] )