From 7eb82e26273cc850c647c4c37903e043b231c064 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 30 Apr 2023 12:21:32 -0400 Subject: [PATCH] remove circular imports from installer --- installer.py | 24 ++++++++---------------- launch.py | 1 - modules/cmd_args.py | 3 +-- webui.py | 2 +- 4 files changed, 10 insertions(+), 20 deletions(-) diff --git a/installer.py b/installer.py index 3be52a6c2..8be8980e1 100644 --- a/installer.py +++ b/installer.py @@ -55,7 +55,6 @@ def setup_logging(clean=False): # check if package is installed def installed(package, friendly: str = None): import pkg_resources - from modules import shared ok = True try: if friendly: @@ -76,7 +75,7 @@ def installed(package, friendly: str = None): ok = ok and spec is not None if ok: version = pkg_resources.get_distribution(p[0]).version - if shared.cmd_opts.use_ipex and p[0] == "pytorch_lightning": + if args.use_ipex and p[0] == "pytorch_lightning": p[1] = "1.8.6" log.debug(f"Package version found: {p[0]} {version}") if len(p) > 1: @@ -93,8 +92,7 @@ def installed(package, friendly: str = None): # install package using pip if not already installed def install(package, friendly: str = None, ignore: bool = False): - from modules import shared - if shared.cmd_opts.use_ipex and package == "pytorch_lightning==1.9.4": + if args.use_ipex and package == "pytorch_lightning==1.9.4": package = "pytorch_lightning==1.8.6" def pip(arg: str): arg = arg.replace('>=', '==') @@ -192,7 +190,6 @@ def check_python(): # check torch version def check_torch(): - from modules import shared if shutil.which('nvidia-smi') is not None or os.path.exists(os.path.join(os.environ.get('SystemRoot') or r'C:\Windows', 'System32', 'nvidia-smi.exe')): log.info('nVidia toolkit detected') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchaudio torchvision --index-url https://download.pytorch.org/whl/cu118') @@ -202,8 +199,7 @@ def check_torch(): os.environ.setdefault('HSA_OVERRIDE_GFX_VERSION', '10.3.0') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision torchaudio --index-url https://download.pytorch.org/whl/rocm5.4.2') xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') - elif shutil.which('sycl-ls') is not None or os.path.exists('/opt/intel/oneapi'): - shared.cmd_opts.use_ipex = True + elif shutil.which('sycl-ls') is not None or os.path.exists('/opt/intel/oneapi') or args.use_ipex: log.info('Intel toolkit detected') torch_command = os.environ.get('TORCH_COMMAND', 'torch==1.13.0a0+git6c9b55e torchvision==0.14.1a0 intel_extension_for_pytorch==1.13.120+xpu --index-url https://developer.intel.com/ipex-whl-stable-xpu') xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') @@ -222,7 +218,7 @@ def check_torch(): try: import torch log.info(f'Torch {torch.__version__}') - if shared.cmd_opts.use_ipex: + if args.use_ipex: import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import log.info(f'Torch backend: Intel OneAPI {torch.__version__}') log.info(f'Torch detected GPU: {torch.xpu.get_device_name("xpu")} VRAM {round(torch.xpu.get_device_properties("xpu").total_memory / 1024 / 1024)}') @@ -373,15 +369,11 @@ def install_submodules(): log.error(f'Error updating submodule: {submodule}') -def ensure_package(pkg): - try: - import pkg # type: ignore - except ImportError: - install(pkg) - - def ensure_base_requirements(): - ensure_package('rich') + try: + import rich # pylint: disable=unused-import + except ImportError: + install('rich', 'rich') def install_requirements(): diff --git a/launch.py b/launch.py index 32e2f8419..76745b4c2 100644 --- a/launch.py +++ b/launch.py @@ -17,7 +17,6 @@ installer.parse_args() import modules.cmd_args args, _ = modules.cmd_args.parser.parse_known_args() - import modules.paths_internal script_path = modules.paths_internal.script_path extensions_dir = modules.paths_internal.extensions_dir diff --git a/modules/cmd_args.py b/modules/cmd_args.py index d365c067a..4f74e5ad0 100644 --- a/modules/cmd_args.py +++ b/modules/cmd_args.py @@ -95,7 +95,6 @@ def compatibility_args(opts, args): opts.dimensions_and_batch_together = True group.add_argument("--lora-dir", help=argparse.SUPPRESS, default=opts.lora_dir) + group.add_argument("--lyco-dir", help=argparse.SUPPRESS, default=opts.lyco_dir) args = parser.parse_args() - if 'lyco_dir' in args: # pylint disable=unsupported-membership-test - args.lyco_dir = opts.lyco_dir return args diff --git a/webui.py b/webui.py index 707585ca7..8d9e0dc2f 100644 --- a/webui.py +++ b/webui.py @@ -13,7 +13,7 @@ startup_timer = timer.Timer() import torch # pylint: disable=C0411 try: - import intel_extension_for_pytorch as ipex + import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import except: pass import torchvision # pylint: disable=W0611,C0411