Update ui
This commit is contained in:
@@ -81,6 +81,7 @@ class Fun_Controller:
|
||||
self.diffusion_transformer_dropdown = model_name
|
||||
self.scheduler_dict = scheduler_dict
|
||||
self.model_type = model_type
|
||||
self.config_path = os.path.realpath(config_path)
|
||||
if config_path is not None:
|
||||
self.config = OmegaConf.load(config_path)
|
||||
self.ulysses_degree = ulysses_degree
|
||||
@@ -94,6 +95,7 @@ class Fun_Controller:
|
||||
self.diffusion_transformer_list = []
|
||||
self.motion_module_list = []
|
||||
self.personalized_model_list = []
|
||||
self.config_list = []
|
||||
|
||||
# config models
|
||||
self.tokenizer = None
|
||||
@@ -107,11 +109,21 @@ class Fun_Controller:
|
||||
self.lora_model_path = "none"
|
||||
self.lora_model_2_path = "none"
|
||||
|
||||
self.refresh_config()
|
||||
self.refresh_diffusion_transformer()
|
||||
self.refresh_personalized_model()
|
||||
if model_name != None:
|
||||
self.update_diffusion_transformer(model_name)
|
||||
|
||||
def refresh_config(self):
|
||||
config_list = []
|
||||
for root, dirs, files in os.walk(self.config_dir):
|
||||
for file in files:
|
||||
if file.endswith(('.yaml', '.yml')):
|
||||
full_path = os.path.join(root, file)
|
||||
config_list.append(full_path)
|
||||
self.config_list = config_list
|
||||
|
||||
def refresh_diffusion_transformer(self):
|
||||
self.diffusion_transformer_list = sorted(glob(os.path.join(self.diffusion_transformer_dir, "*/")))
|
||||
|
||||
@@ -122,6 +134,11 @@ class Fun_Controller:
|
||||
def update_model_type(self, model_type):
|
||||
self.model_type = model_type
|
||||
|
||||
def update_config(self, config_dropdown):
|
||||
self.config_path = config_dropdown
|
||||
self.config = OmegaConf.load(config_dropdown)
|
||||
print(f"Update config: {config_dropdown}")
|
||||
|
||||
def update_diffusion_transformer(self, diffusion_transformer_dropdown):
|
||||
pass
|
||||
|
||||
|
||||
@@ -336,3 +336,23 @@ def create_ui_outputs():
|
||||
interactive=False
|
||||
)
|
||||
return result_image, result_video, infer_progress
|
||||
|
||||
def create_config(controller):
|
||||
gr.Markdown(
|
||||
"""
|
||||
### Config Path (配置文件路径)
|
||||
"""
|
||||
)
|
||||
with gr.Row():
|
||||
config_dropdown = gr.Dropdown(
|
||||
label="Config Path (配置文件路径)",
|
||||
choices=controller.config_list,
|
||||
value=controller.config_path,
|
||||
interactive=True,
|
||||
)
|
||||
config_refresh_button = gr.Button(value="\U0001F503", elem_classes="toolbutton")
|
||||
def refresh_config():
|
||||
controller.refresh_config()
|
||||
return gr.update(choices=controller.config_list)
|
||||
config_refresh_button.click(fn=refresh_config, inputs=[], outputs=[config_dropdown])
|
||||
return config_dropdown, config_refresh_button
|
||||
@@ -35,7 +35,7 @@ from .ui import (create_cfg_and_seedbox, create_cfg_riflex_k,
|
||||
create_generation_methods_and_video_length,
|
||||
create_height_width, create_model_checkpoints,
|
||||
create_model_type, create_prompts, create_samplers,
|
||||
create_teacache_params, create_ui_outputs)
|
||||
create_teacache_params, create_ui_outputs, create_config)
|
||||
from ..dist import set_multi_gpus_devices, shard_model
|
||||
|
||||
|
||||
@@ -388,6 +388,7 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype,
|
||||
"""
|
||||
)
|
||||
with gr.Column(variant="panel"):
|
||||
config_dropdown, config_refresh_button = create_config(controller)
|
||||
model_type = create_model_type(visible=False)
|
||||
diffusion_transformer_dropdown, diffusion_transformer_refresh_button = \
|
||||
create_model_checkpoints(controller, visible=True)
|
||||
@@ -428,6 +429,12 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, compile_dit, weight_dtype,
|
||||
|
||||
result_image, result_video, infer_progress = create_ui_outputs()
|
||||
|
||||
config_dropdown.change(
|
||||
fn=controller.update_config,
|
||||
inputs=[config_dropdown],
|
||||
outputs=[]
|
||||
)
|
||||
|
||||
model_type.change(
|
||||
fn=controller.update_model_type,
|
||||
inputs=[model_type],
|
||||
|
||||
Reference in New Issue
Block a user