diff --git a/installer.py b/installer.py index d16acc70d..53027f5b6 100644 --- a/installer.py +++ b/installer.py @@ -641,14 +641,16 @@ def install_rocm_zluda(): else: #check_python(supported_minors=[10, 11, 12, 13, 14], reason='ROCm backend requires a Python version between 3.10 and 3.13') if args.use_nightly: - if rocm.version is None or float(rocm.version) >= 7.1: # assume the latest if version check fails + if rocm.version is None or float(rocm.version) >= 7.2: # assume the latest if version check fails + torch_command = os.environ.get('TORCH_COMMAND', '--upgrade --pre torch torchvision --index-url https://download.pytorch.org/whl/nightly/rocm7.2') + else: # oldest rocm version on nightly is 7.1 torch_command = os.environ.get('TORCH_COMMAND', '--upgrade --pre torch torchvision --index-url https://download.pytorch.org/whl/nightly/rocm7.1') - else: # oldest rocm version on nightly is 7.0 - torch_command = os.environ.get('TORCH_COMMAND', '--upgrade --pre torch torchvision --index-url https://download.pytorch.org/whl/nightly/rocm7.0') else: - if rocm.version is None or float(rocm.version) >= 7.1: # assume the latest if version check fails - torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.10.0+rocm7.1 torchvision==0.25.0+rocm7.1 --index-url https://download.pytorch.org/whl/rocm7.1') - elif rocm.version == "7.0": # assume the latest if version check fails + if rocm.version is None or float(rocm.version) >= 7.2: # assume the latest if version check fails + torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.11.0+rocm7.2 torchvision==0.26.0+rocm7.2 --index-url https://download.pytorch.org/whl/rocm7.2') + elif rocm.version == "7.1": + torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.11.0+rocm7.1 torchvision==0.26.0+rocm7.1 --index-url https://download.pytorch.org/whl/rocm7.1') + elif rocm.version == "7.0": torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.10.0+rocm7.0 torchvision==0.25.0+rocm7.0 --index-url https://download.pytorch.org/whl/rocm7.0') elif rocm.version == "6.4": torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.9.1+rocm6.4 torchvision==0.24.1+rocm6.4 --index-url https://download.pytorch.org/whl/rocm6.4')