From 2bc31d53ed9c1d449cdae506fba9e0d9f7b6596d Mon Sep 17 00:00:00 2001 From: Disty0 Date: Thu, 5 Sep 2024 15:48:42 +0300 Subject: [PATCH] ROCm fix installer getting stuck on onnxruntime --- installer.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/installer.py b/installer.py index 413887ff3..6b08b6ad9 100644 --- a/installer.py +++ b/installer.py @@ -536,9 +536,13 @@ def install_rocm_zluda(): else: torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://download.pytorch.org/whl/rocm{rocm.version}') - ort_version = os.environ.get('ONNXRUNTIME_VERSION', None) - ort_package = os.environ.get('ONNXRUNTIME_PACKAGE', f"--pre onnxruntime-training{'' if ort_version is None else ('==' + ort_version)} --index-url https://pypi.lsh.sh/{rocm.version[0]}{rocm.version[2]} --extra-index-url https://pypi.org/simple") - install(ort_package, 'onnxruntime-training') + if sys.version_info < (3, 11): + ort_version = os.environ.get('ONNXRUNTIME_VERSION', None) + if rocm.version is None or float(rocm.version) > 6.0: + ort_package = os.environ.get('ONNXRUNTIME_PACKAGE', f"--pre onnxruntime-training{'' if ort_version is None else ('==' + ort_version)} --index-url https://pypi.lsh.sh/60 --extra-index-url https://pypi.org/simple") + else: + ort_package = os.environ.get('ONNXRUNTIME_PACKAGE', f"--pre onnxruntime-training{'' if ort_version is None else ('==' + ort_version)} --index-url https://pypi.lsh.sh/{rocm.version[0]}{rocm.version[2]} --extra-index-url https://pypi.org/simple") + install(ort_package, 'onnxruntime-training') if hip_default_device is not None and rocm.version != "6.2" and rocm.version == rocm.version_torch and rocm.get_blaslt_enabled(): log.debug(f'hipBLASLt arch={hip_default_device.name} available={hip_default_device.blaslt_supported}')