From 765aeb94fd008e4971e015d9af57a8d6a8e1a7e3 Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Sun, 8 Oct 2023 14:44:35 +0000 Subject: [PATCH 1/7] Update to PyTorch 2.1.0 with pre-release xformers. --- installer.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/installer.py b/installer.py index 4da612567..36e783fec 100644 --- a/installer.py +++ b/installer.py @@ -361,7 +361,7 @@ def check_torch(): elif allow_cuda and (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 CUDA toolkit detected') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu121') - xformers_package = os.environ.get('XFORMERS_PACKAGE', 'xformers==0.0.22' if opts.get('cross_attention_optimization', '') == 'xFormers' else 'none') + xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers<0.0.24' if opts.get('cross_attention_optimization', '') == 'xFormers' else 'none') elif allow_rocm and (shutil.which('rocminfo') is not None or os.path.exists('/opt/rocm/bin/rocminfo') or os.path.exists('/dev/kfd')): log.info('AMD ROCm toolkit detected') os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512') @@ -397,18 +397,18 @@ def check_torch(): log.debug(f'HSA_OVERRIDE_GFX_VERSION auto config is skipped for {gpu}') try: command = subprocess.run('hipconfig --version', shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE) - arr = command.stdout.decode(encoding="utf8", errors="ignore").split('.') - if len(arr) >= 2: - rocm_ver = f'{arr[0]}.{arr[1]}' + rocm_ver_tuple = tuple(int(v) for v in command.stdout.decode(encoding="utf8", errors="ignore").split('.')) + if len(rocm_ver_tuple) >= 2: + rocm_ver = f'{rocm_ver_tuple[0]}.{rocm_ver_tuple[1]}' log.debug(f'ROCm version detected: {rocm_ver}') except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None - if rocm_ver in ['5.5', '5.6', '5.7']: + if rocm_ver_tuple in {(5, 7)}: # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday torch_command = os.environ.get('TORCH_COMMAND', f'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') else: - torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.1 torchvision==0.15.2 --index-url https://download.pytorch.org/whl/rocm5.4.2') + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm{rocm_ver}') xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') elif allow_ipex and (args.use_ipex or shutil.which('sycl-ls') is not None or shutil.which('sycl-ls.exe') is not None or os.environ.get('ONEAPI_ROOT') is not None or os.path.exists('/opt/intel/oneapi') or os.path.exists("C:/Program Files (x86)/Intel/oneAPI") or os.path.exists("C:/oneAPI")): args.use_ipex = True # pylint: disable=attribute-defined-outside-init @@ -430,7 +430,7 @@ def check_torch(): else: machine = platform.machine() if sys.platform == 'darwin': - torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.1 torchvision==0.15.2') + torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision') elif allow_directml and args.use_directml and ('arm' not in machine and 'aarch' not in machine): log.info('Using DirectML Backend') torch_command = os.environ.get('TORCH_COMMAND', 'torch-directml') @@ -700,7 +700,7 @@ def ensure_base_requirements(): except ImportError: install('setuptools', 'setuptools') try: - import setuptools # pylint: disable=unused-import + import setuptools # pylint: disable=unused-import # noqa: F811 except ImportError: pass try: @@ -708,7 +708,7 @@ def ensure_base_requirements(): except ImportError: install('rich', 'rich') try: - import rich # pylint: disable=unused-import + import rich # pylint: disable=unused-import # noqa: F811 except ImportError: pass From d4a728ccaa8fecf59481cc92e00f51801e4da038 Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Sun, 8 Oct 2023 14:57:34 +0000 Subject: [PATCH 2/7] Install right PyTorch version if xformers is enabled. --- installer.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/installer.py b/installer.py index 36e783fec..96ed6276c 100644 --- a/installer.py +++ b/installer.py @@ -360,8 +360,10 @@ def check_torch(): pass elif allow_cuda and (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 CUDA toolkit detected') - torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu121') - xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers<0.0.24' if opts.get('cross_attention_optimization', '') == 'xFormers' else 'none') + xformers_enabled = opts.get('cross_attention_optimization', '') == 'xFormers' + cuda_version = "118" if xformers_enabled else "121" + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/cu{cuda_version}') + xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers<0.0.24' if xformers_enabled else 'none') elif allow_rocm and (shutil.which('rocminfo') is not None or os.path.exists('/opt/rocm/bin/rocminfo') or os.path.exists('/dev/kfd')): log.info('AMD ROCm toolkit detected') os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512') From 8d35d681670bcf60c8f90c15b7c4327c80998fec Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Sun, 8 Oct 2023 18:33:25 +0200 Subject: [PATCH 3/7] Revert some changes. --- installer.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/installer.py b/installer.py index 96ed6276c..9f255192e 100644 --- a/installer.py +++ b/installer.py @@ -360,10 +360,8 @@ def check_torch(): pass elif allow_cuda and (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 CUDA toolkit detected') - xformers_enabled = opts.get('cross_attention_optimization', '') == 'xFormers' - cuda_version = "118" if xformers_enabled else "121" - torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/cu{cuda_version}') - xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers<0.0.24' if xformers_enabled else 'none') + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/cu121') + xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers<0.0.24' if opts.get('cross_attention_optimization', '') == 'xFormers' else 'none') elif allow_rocm and (shutil.which('rocminfo') is not None or os.path.exists('/opt/rocm/bin/rocminfo') or os.path.exists('/dev/kfd')): log.info('AMD ROCm toolkit detected') os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512') @@ -406,11 +404,12 @@ def check_torch(): except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None - if rocm_ver_tuple in {(5, 7)}: + if rocm_ver_tuple[:2] <= (5, 6): # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday torch_command = os.environ.get('TORCH_COMMAND', f'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') else: - torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm{rocm_ver}') + # ROCm 5.5 is oldest for PyTorch 2.1 + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm5.5') xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') elif allow_ipex and (args.use_ipex or shutil.which('sycl-ls') is not None or shutil.which('sycl-ls.exe') is not None or os.environ.get('ONEAPI_ROOT') is not None or os.path.exists('/opt/intel/oneapi') or os.path.exists("C:/Program Files (x86)/Intel/oneAPI") or os.path.exists("C:/oneAPI")): args.use_ipex = True # pylint: disable=attribute-defined-outside-init From 0bf43ec2cf567c684ee089246816af21652fc918 Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Sun, 8 Oct 2023 18:34:48 +0200 Subject: [PATCH 4/7] Add comment. --- installer.py | 1 + 1 file changed, 1 insertion(+) diff --git a/installer.py b/installer.py index 9f255192e..f700b63e2 100644 --- a/installer.py +++ b/installer.py @@ -404,6 +404,7 @@ def check_torch(): except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None + # ROCm 5.7 and above don't have PyTorch 2.1.0. if rocm_ver_tuple[:2] <= (5, 6): # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday torch_command = os.environ.get('TORCH_COMMAND', f'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') From 498c23824f27ed5826e891e8b39164e9c3709383 Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Thu, 12 Oct 2023 14:40:19 +0200 Subject: [PATCH 5/7] Revert tuple compoarison. --- installer.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/installer.py b/installer.py index c7545f98c..9ad85ed6b 100644 --- a/installer.py +++ b/installer.py @@ -405,10 +405,9 @@ def check_torch(): except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None - # ROCm 5.7 and above don't have PyTorch 2.1.0. - if rocm_ver_tuple[:2] <= (5, 6): + if rocm_ver_tuple[:2] == "5.7": # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday - torch_command = os.environ.get('TORCH_COMMAND', f'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') + torch_command = os.environ.get('TORCH_COMMAND', 'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') else: # ROCm 5.5 is oldest for PyTorch 2.1 torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm5.5') From 9e227e24e14d6a78bda1f994b6b84a086141a61b Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Thu, 12 Oct 2023 15:48:00 +0200 Subject: [PATCH 6/7] Revert tuple changes. --- installer.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/installer.py b/installer.py index 9ad85ed6b..9d4ffbdc8 100644 --- a/installer.py +++ b/installer.py @@ -398,14 +398,14 @@ def check_torch(): log.debug(f'HSA_OVERRIDE_GFX_VERSION auto config is skipped for {gpu}') try: command = subprocess.run('hipconfig --version', shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE) - rocm_ver_tuple = tuple(int(v) for v in command.stdout.decode(encoding="utf8", errors="ignore").split('.')) - if len(rocm_ver_tuple) >= 2: - rocm_ver = f'{rocm_ver_tuple[0]}.{rocm_ver_tuple[1]}' + arr = command.stdout.decode(encoding="utf8", errors="ignore").split('.') + if len(arr) >= 2: + rocm_ver = f'{arr[0]}.{arr[1]}' log.debug(f'ROCm version detected: {rocm_ver}') except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None - if rocm_ver_tuple[:2] == "5.7": + if rocm_ver in {"5.5", "5.6"}: # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday torch_command = os.environ.get('TORCH_COMMAND', 'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') else: From d07fae0baf6baa8b8f1370fa0f361cdee03e620f Mon Sep 17 00:00:00 2001 From: Hameer Abbasi Date: Thu, 12 Oct 2023 15:50:49 +0200 Subject: [PATCH 7/7] Better detection of supported ROCm for PyTorch 2.1.0. --- installer.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/installer.py b/installer.py index 9d4ffbdc8..b128e4fed 100644 --- a/installer.py +++ b/installer.py @@ -405,9 +405,11 @@ def check_torch(): except Exception as e: log.debug(f'ROCm hipconfig failed: {e}') rocm_ver = None - if rocm_ver in {"5.5", "5.6"}: + if rocm_ver in {"5.7"}: # install torch nightly via torchvision to avoid wasting bandwidth when torchvision depends on torch from yesterday - torch_command = os.environ.get('TORCH_COMMAND', 'torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') + torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --pre --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') + elif rocm_ver in {"5.5", "5.6"}: + torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/nightly/rocm{rocm_ver}') else: # ROCm 5.5 is oldest for PyTorch 2.1 torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm5.5')