Update ui

This commit is contained in:
bubbliiiing
2025-07-30 14:20:47 +08:00
parent 48c8288323
commit e258d4158b
3 changed files with 45 additions and 1 deletions
+17
View File
@@ -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
+20
View File
@@ -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
+8 -1
View File
@@ -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],