mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
minor updates
This commit is contained in:
@@ -197,10 +197,11 @@ def initialize_onnx():
|
||||
from installer import log
|
||||
from modules import devices
|
||||
from modules.shared import opts
|
||||
from .execution_providers import ExecutionProvider, TORCH_DEVICE_TO_EP, available_execution_providers
|
||||
if devices.backend == "rocm":
|
||||
TORCH_DEVICE_TO_EP["cuda"] = ExecutionProvider.ROCm
|
||||
try:
|
||||
try: # may fail on onnx import
|
||||
import onnx # pylint: disable=unused-import
|
||||
from .execution_providers import ExecutionProvider, TORCH_DEVICE_TO_EP, available_execution_providers
|
||||
if devices.backend == "rocm":
|
||||
TORCH_DEVICE_TO_EP["cuda"] = ExecutionProvider.ROCm
|
||||
from .pipelines.onnx_stable_diffusion_pipeline import OnnxStableDiffusionPipeline
|
||||
from .pipelines.onnx_stable_diffusion_img2img_pipeline import OnnxStableDiffusionImg2ImgPipeline
|
||||
from .pipelines.onnx_stable_diffusion_inpaint_pipeline import OnnxStableDiffusionInpaintPipeline
|
||||
@@ -208,8 +209,7 @@ def initialize_onnx():
|
||||
from .pipelines.onnx_stable_diffusion_xl_pipeline import OnnxStableDiffusionXLPipeline
|
||||
from .pipelines.onnx_stable_diffusion_xl_img2img_pipeline import OnnxStableDiffusionXLImg2ImgPipeline
|
||||
|
||||
# OnnxRuntimeModel Hijack.
|
||||
OnnxRuntimeModel.__module__ = 'diffusers'
|
||||
OnnxRuntimeModel.__module__ = 'diffusers' # OnnxRuntimeModel Hijack.
|
||||
diffusers.OnnxRuntimeModel = OnnxRuntimeModel
|
||||
|
||||
diffusers.OnnxStableDiffusionPipeline = OnnxStableDiffusionPipeline
|
||||
@@ -229,8 +229,7 @@ def initialize_onnx():
|
||||
diffusers.OnnxStableDiffusionXLImg2ImgPipeline = OnnxStableDiffusionXLImg2ImgPipeline
|
||||
diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["onnx-stable-diffusion-xl"] = diffusers.OnnxStableDiffusionXLImg2ImgPipeline
|
||||
|
||||
# Huggingface model compatibility
|
||||
diffusers.ORTStableDiffusionXLPipeline = diffusers.OnnxStableDiffusionXLPipeline
|
||||
diffusers.ORTStableDiffusionXLPipeline = diffusers.OnnxStableDiffusionXLPipeline # Huggingface model compatibility
|
||||
diffusers.ORTStableDiffusionXLImg2ImgPipeline = diffusers.OnnxStableDiffusionXLImg2ImgPipeline
|
||||
|
||||
optimum.onnxruntime.modeling_diffusion._ORTDiffusionModelPart.to = ORTDiffusionModelPart_to # pylint: disable=protected-access
|
||||
|
||||
@@ -124,9 +124,7 @@ 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", encoding="utf-8") as file:
|
||||
config = json.load(file)
|
||||
sd_configs[submodel] = config
|
||||
@@ -142,26 +140,21 @@ def create_ui():
|
||||
|
||||
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(): # pylint: disable=protected-access
|
||||
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
|
||||
@@ -192,9 +185,7 @@ 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", encoding="utf-8") as file:
|
||||
config = json.load(file)
|
||||
sdxl_configs[submodel] = config
|
||||
@@ -210,31 +201,25 @@ def create_ui():
|
||||
|
||||
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): # pylint: disable=function-redefined
|
||||
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(): # pylint: disable=protected-access
|
||||
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():
|
||||
|
||||
Reference in New Issue
Block a user