Update api and ui
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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():
|
||||
|
||||
@@ -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]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user