From d2a481bc7579cb19e546528a185640799a5da6fa Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Apr 2023 09:32:32 -0400 Subject: [PATCH] update setup --- requirements.txt | 2 +- setup.py | 10 +++++++--- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/requirements.txt b/requirements.txt index 80c495d1e..d4b437458 100644 --- a/requirements.txt +++ b/requirements.txt @@ -62,4 +62,4 @@ pytorch_lightning==1.9.4 tensorflow==2.12.0 transformers==4.26.1 timm==0.6.13 -tomesd>=0.1.2 +tomesd==0.1.2 diff --git a/setup.py b/setup.py index 4031d6db7..f11057c83 100644 --- a/setup.py +++ b/setup.py @@ -56,12 +56,15 @@ def installed(package, friendly: str = None): ok = True try: if friendly: - pkgs = [friendly] + pkgs = friendly.split() else: pkgs = [p for p in package.split() if not p.startswith('-') and not p.startswith('=')] pkgs = [p.split('/')[-1] for p in pkgs] # get only package name if installing from url for pkg in pkgs: - p = pkg.split('==') + if '>=' in pkg: + p = pkg.split('>=') + else: + p = pkg.split('==') spec = pkg_resources.working_set.by_key.get(p[0], None) # more reliable than importlib if spec is None: spec = pkg_resources.working_set.by_key.get(p[0].lower(), None) # check name variations @@ -86,6 +89,7 @@ def installed(package, friendly: str = None): # install package using pip if not already installed def install(package, friendly: str = None, ignore: bool = False): def pip(arg: str): + arg = arg.replace('>=', '==') log.info(f'Installing package: {arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace(" ", " ").strip()}') log.debug(f"Running pip: {arg}") result = subprocess.run(f'"{sys.executable}" -m pip {arg}', shell=True, check=False, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE) @@ -180,7 +184,7 @@ def check_torch(): torch_command = os.environ.get('TORCH_COMMAND', 'torch torchaudio torchvision') xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') if 'torch' in torch_command: - install(torch_command) + install(torch_command, 'torch torchvision torchaudio') try: import torch log.info(f'Torch {torch.__version__}')