Update api and ui

This commit is contained in:
bubbliiiing
2025-07-30 12:53:45 +08:00
parent 1aacfe6bca
commit 48c8288323
5 changed files with 77 additions and 16 deletions
+4
View File
@@ -93,7 +93,9 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller):
datas: dict,
):
base_model_path = datas.get('base_model_path', 'none')
base_model_2_path = datas.get('base_model_2_path', 'none')
lora_model_path = datas.get('lora_model_path', 'none')
lora_model_2_path = datas.get('lora_model_2_path', 'none')
lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)
prompt_textbox = datas.get('prompt_textbox', None)
negative_prompt_textbox = datas.get('negative_prompt_textbox', 'The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ')
@@ -205,6 +207,8 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller):
cfg_skip_ratio = cfg_skip_ratio,
enable_riflex = enable_riflex,
riflex_k = riflex_k,
base_model_2_path = base_model_2_path,
lora_model_2_path = lora_model_2_path,
fps = fps,
is_api = True,
)
+4
View File
@@ -99,7 +99,9 @@ if ray is not None:
def generate(self, datas):
try:
base_model_path = datas.get('base_model_path', 'none')
base_model_2_path = datas.get('base_model_2_path', 'none')
lora_model_path = datas.get('lora_model_path', 'none')
lora_model_2_path = datas.get('lora_model_2_path', 'none')
lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)
prompt_textbox = datas.get('prompt_textbox', None)
negative_prompt_textbox = datas.get('negative_prompt_textbox', 'The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ')
@@ -211,6 +213,8 @@ if ray is not None:
cfg_skip_ratio = cfg_skip_ratio,
enable_riflex = enable_riflex,
riflex_k = riflex_k,
base_model_2_path = base_model_2_path,
lora_model_2_path = lora_model_2_path,
fps = fps,
is_api = True,
)
+22 -6
View File
@@ -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,):
+20 -2
View File
@@ -79,7 +79,7 @@ def create_fake_model_checkpoints(model_name, visible):
)
return diffusion_transformer_dropdown
def create_finetune_models_checkpoints(controller, visible):
def create_finetune_models_checkpoints(controller, visible, add_checkpoint_2=False):
with gr.Row(visible=visible):
base_model_dropdown = gr.Dropdown(
label="Select base Dreambooth model (选择基模型[非必需])",
@@ -87,6 +87,13 @@ def create_finetune_models_checkpoints(controller, visible):
value="none",
interactive=True,
)
if add_checkpoint_2:
base_model_2_dropdown = gr.Dropdown(
label="Select base Dreambooth model (选择第二个基模型[非必需])",
choices=["none"] + controller.personalized_model_list,
value="none",
interactive=True,
)
lora_model_dropdown = gr.Dropdown(
label="Select LoRA model (选择LoRA模型[非必需])",
@@ -94,6 +101,13 @@ def create_finetune_models_checkpoints(controller, visible):
value="none",
interactive=True,
)
if add_checkpoint_2:
lora_model_2_dropdown = gr.Dropdown(
label="Select LoRA model (选择LoRA模型[非必需])",
choices=["none"] + controller.personalized_model_list,
value="none",
interactive=True,
)
lora_alpha_slider = gr.Slider(label="LoRA alpha (LoRA权重)", value=0.55, minimum=0, maximum=2, interactive=True)
@@ -106,7 +120,11 @@ def create_finetune_models_checkpoints(controller, visible):
]
personalized_refresh_button.click(fn=update_personalized_model, inputs=[], outputs=[base_model_dropdown, lora_model_dropdown])
return base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button
if not add_checkpoint_2:
return base_model_dropdown, lora_model_dropdown, lora_alpha_slider, personalized_refresh_button
else:
return [base_model_dropdown, base_model_2_dropdown], [lora_model_dropdown, lora_model_2_dropdown], \
lora_alpha_slider, personalized_refresh_button
def create_fake_finetune_models_checkpoints(visible):
with gr.Row():
+27 -8
View File
@@ -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]
)