From fcf00bd8549cae539c0cef4b4bf57842f8779b8a Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Fri, 2 Feb 2024 05:20:46 +0900 Subject: [PATCH] move initialization of onnx pipelines & create onnx folder automatically --- modules/onnx_impl/__init__.py | 6 ++++++ modules/shared.py | 2 ++ modules/shared_items.py | 3 --- 3 files changed, 8 insertions(+), 3 deletions(-) diff --git a/modules/onnx_impl/__init__.py b/modules/onnx_impl/__init__.py index 8e5ef7121..8b8fdef14 100644 --- a/modules/onnx_impl/__init__.py +++ b/modules/onnx_impl/__init__.py @@ -1,3 +1,4 @@ +import os from typing import Any, Dict, Optional import torch import diffusers @@ -142,9 +143,14 @@ def initialize(): return from modules import devices + from modules.paths import models_path from . import pipelines from .execution_providers import ExecutionProvider, TORCH_DEVICE_TO_EP + onnx_dir = os.path.join(models_path, "ONNX") + if not os.path.isdir(onnx_dir): + os.mkdir(onnx_dir) + if devices.backend == "rocm": TORCH_DEVICE_TO_EP["cuda"] = ExecutionProvider.ROCm diff --git a/modules/shared.py b/modules/shared.py index d3454dddf..b7389b499 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -16,6 +16,7 @@ from rich.console import Console from modules import errors, shared_items, shared_state, cmd_args, theme from modules.paths import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611 from modules.dml import memory_providers, default_memory_provider, directml_do_hijack +from modules.onnx_impl import initialize as initialize_onnx from modules.onnx_impl.execution_providers import available_execution_providers, get_default_execution_provider import modules.interrogate import modules.memmon @@ -906,6 +907,7 @@ mem_mon = modules.memmon.MemUsageMonitor("MemMon", devices.device) max_workers = 4 if devices.backend == "directml": directml_do_hijack() +initialize_onnx() class TotalTQDM: # compatibility with previous global-tqdm diff --git a/modules/shared_items.py b/modules/shared_items.py index 71487bfe0..3d1240c43 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -27,9 +27,6 @@ def list_crossattention(): def get_pipelines(): import diffusers from installer import log - from modules.onnx_impl import initialize as initialize_onnx_pipelines - - initialize_onnx_pipelines() pipelines = { # note: not all pipelines can be used manually as they require prior pipeline next to decoder pipeline 'Autodetect': None,