Implement Olive config editor.

This commit is contained in:
Seunghoon Lee
2024-01-14 03:04:59 +09:00
parent f5a7a53dbd
commit d55d92bab4
+78 -15
View File
@@ -68,7 +68,10 @@ def create_ui():
onnx_vae_apply_button.click(fn=onnx_vae_apply, inputs=[onnx_vae_id, onnx_vae_subfolder,])
if opts.cuda_compile_backend == "olive-ai":
from olive.passes import REGISTRY
import olive.passes as olive_passes
from olive.hardware.accelerator import AcceleratorSpec, Device
accelerator = AcceleratorSpec(accelerator_type=Device.GPU, execution_provider=opts.onnx_execution_provider)
with gr.Tabs(elem_id="tabs_olive"):
with gr.TabItem("Manage cache", id="manage_cache"):
@@ -100,10 +103,12 @@ def create_ui():
shutil.rmtree(os.path.join(opts.onnx_cached_models_path, f"{dirname}-{size[0]}w-{size[1]}h"))
log.info(f"Olive processed cache of '{dirname}' is removed: width={size[0]}, height={size[1]}")
cache_list_optimized_headers = ["height", "width"]
cache_list_optimized_types = ["str", "str"]
cache_list_optimized = gr.DataFrame(None, label="Optimized caches", show_label=True, overflow_row_behaviour='paginate', interactive=False, max_rows=10, headers=cache_list_optimized_headers, datatype=cache_list_optimized_types, type="array")
cache_list_optimized.select(fn=select_cache_optimized, inputs=[cache_list_optimized,], outputs=[cache_optimized_selected,])
with gr.Row():
cache_list_optimized_headers = ["height", "width"]
cache_list_optimized_types = ["str", "str"]
cache_list_optimized = gr.Dataframe(None, label="Optimized caches", show_label=True, overflow_row_behaviour='paginate', interactive=False, max_rows=10, headers=cache_list_optimized_headers, datatype=cache_list_optimized_types, type="array")
cache_list_optimized.select(fn=select_cache_optimized, inputs=[cache_list_optimized,], outputs=[cache_optimized_selected,])
cache_remove_optimized = gr.Button(value="Remove selected cache", visible=False)
cache_remove_optimized.click(fn=remove_cache_optimized, inputs=[cache_state_dirname, cache_optimized_selected,])
@@ -146,7 +151,8 @@ def create_ui():
with gr.TabItem("Stable Diffusion", id="sd"):
sd_config_path = os.path.join(sd_configs_path, "olive", "sd")
sd_submodels = os.listdir(sd_config_path)
sd_configs: Dict[str, Dict] = {}
sd_configs: Dict[str, Dict[str, Dict[str, Dict]]] = {}
sd_pass_config_components: Dict[str, Dict[str, Dict]] = {}
with gr.Tabs(elem_id="tabs_sd_submodel"):
def sd_create_change_listener(*args):
@@ -157,6 +163,8 @@ def create_ui():
for submodel in sd_submodels:
config: Dict = None
sd_pass_config_components[submodel] = {}
with open(os.path.join(sd_config_path, submodel), "r") as file:
config = json.load(file)
sd_configs[submodel] = config
@@ -167,11 +175,37 @@ def create_ui():
pass_flows.change(fn=sd_create_change_listener(submodel, "pass_flows", 0), inputs=pass_flows)
with gr.Tabs(elem_id=f"tabs_sd_{submodel_name}_pass"):
for k in sd_configs[submodel]["passes"]:
with gr.TabItem(k, id=f"sd_{submodel_name}_pass_{k}"):
pass_type = gr.Dropdown(label="Type", value=sd_configs[submodel]["passes"][k]["type"], choices=(x.__name__ for x in tuple(REGISTRY.values())))
for pass_name in sd_configs[submodel]["passes"]:
sd_pass_config_components[submodel][pass_name] = {}
pass_type.change(fn=sd_create_change_listener(submodel, "passes", k, "type"), inputs=pass_type)
with gr.TabItem(pass_name, id=f"sd_{submodel_name}_pass_{pass_name}"):
config_dict = sd_configs[submodel]["passes"][pass_name]
pass_type = gr.Dropdown(label="Type", value=config_dict["type"], choices=(x.__name__ for x in tuple(olive_passes.REGISTRY.values())))
def create_pass_config_change_listener(submodel, pass_name, config_key):
def listener(value):
sd_configs[submodel]["passes"][pass_name]["config"][config_key] = value
return listener
for config_key, v in getattr(olive_passes, config_dict["type"], olive_passes.Pass)._default_config(accelerator).items():
component = None
if v.type_ == bool:
component = gr.Checkbox
elif v.type_ == str:
component = gr.Textbox
elif v.type_ == int:
component = gr.Number
if component is not None:
component = component(value=config_dict["config"][config_key] if config_key in config_dict["config"] else v.default_value, label=config_key)
sd_pass_config_components[submodel][pass_name][config_key] = component
component.change(fn=create_pass_config_change_listener(submodel, pass_name, config_key), inputs=component)
pass_type.change(fn=sd_create_change_listener(submodel, "passes", config_key, "type"), inputs=pass_type)
def sd_save():
for k, v in sd_configs.items():
@@ -185,7 +219,8 @@ def create_ui():
with gr.TabItem("Stable Diffusion XL", id="sdxl"):
sdxl_config_path = os.path.join(sd_configs_path, "olive", "sdxl")
sdxl_submodels = os.listdir(sdxl_config_path)
sdxl_configs: Dict[str, Dict] = {}
sdxl_configs: Dict[str, Dict[str, Dict[str, Dict]]] = {}
sdxl_pass_config_components: Dict[str, Dict[str, Dict]] = {}
with gr.Tabs(elem_id="tabs_sdxl_submodel"):
def sdxl_create_change_listener(*args):
@@ -196,6 +231,8 @@ def create_ui():
for submodel in sdxl_submodels:
config: Dict = None
sdxl_pass_config_components[submodel] = {}
with open(os.path.join(sdxl_config_path, submodel), "r") as file:
config = json.load(file)
sdxl_configs[submodel] = config
@@ -206,11 +243,37 @@ def create_ui():
pass_flows.change(fn=sdxl_create_change_listener(submodel, "pass_flows", 0), inputs=pass_flows)
with gr.Tabs(elem_id=f"tabs_sdxl_{submodel_name}_pass"):
for k in sdxl_configs[submodel]["passes"]:
with gr.TabItem(k, id=f"sdxl_{submodel_name}_pass_{k}"):
pass_type = gr.Dropdown(label="Type", value=sdxl_configs[submodel]["passes"][k]["type"], choices=(x.__name__ for x in tuple(REGISTRY.values())))
for pass_name in sdxl_configs[submodel]["passes"]:
sdxl_pass_config_components[submodel][pass_name] = {}
pass_type.change(fn=sdxl_create_change_listener(submodel, "passes", k, "type"), inputs=pass_type)
with gr.TabItem(pass_name, id=f"sdxl_{submodel_name}_pass_{pass_name}"):
config_dict = sdxl_configs[submodel]["passes"][pass_name]
pass_type = gr.Dropdown(label="Type", value=sdxl_configs[submodel]["passes"][pass_name]["type"], choices=(x.__name__ for x in tuple(olive_passes.REGISTRY.values())))
def create_pass_config_change_listener(submodel, pass_name, config_key):
def listener(value):
sdxl_configs[submodel]["passes"][pass_name]["config"][config_key] = value
return listener
for config_key, v in getattr(olive_passes, config_dict["type"], olive_passes.Pass)._default_config(accelerator).items():
component = None
if v.type_ == bool:
component = gr.Checkbox
elif v.type_ == str:
component = gr.Textbox
elif v.type_ == int:
component = gr.Number
if component is not None:
component = component(value=config_dict["config"][config_key] if config_key in config_dict["config"] else v.default_value, label=config_key)
sdxl_pass_config_components[submodel][pass_name][config_key] = component
component.change(fn=create_pass_config_change_listener(submodel, pass_name, config_key), inputs=component)
pass_type.change(fn=sdxl_create_change_listener(submodel, "passes", pass_name, "type"), inputs=pass_type)
def sdxl_save():
for k, v in sdxl_configs.items():