PyTorch deep learning के लिए go-to framework है, और AMD hardware पर यह ROCm के ज़रिए चलता है। एक बार ROCm install हो जाए, तो Radeon GPU पर PyTorch चलाना सिर्फ pip install की बात है।

मैंने इसे Radeon RX 6800M (gfx1031, RDNA2) पर ROCm 7.14 के साथ Ubuntu 26.04 पर test किया है – नीचे दिए गए exact commands वही हैं जो काम किए। यही flow दूसरे RDNA2/RDNA3 cards पर भी लागू होता है; बस install command में gfx target बदल दें।

Prerequisites

Step 1 – virtual environment बनाएँ

हमेशा PyTorch को fresh virtual environment में install करें ताकि यह कभी system packages या दूसरे Python projects से conflict न करे:

python3.12 -m venv .venv

आपके पास जो भी Python version हो उसका इस्तेमाल करें – python3.11, python3.13 या python3.14 सभी काम करते हैं। यह आपकी current directory में .venv नाम का folder बनाता है।

Step 2 – environment को activate करें

इसे activate करें ताकि python और pip venv की ओर point करें:

source .venv/bin/activate

आपको पता चलेगा कि यह काम कर गया जब prompt की शुरुआत में (.venv) दिखे।

Step 3 – ROCm support के साथ PyTorch install करें

AMD के wheel repository से ROCm-enabled PyTorch, torchvision और torchaudio install करें। अपने GPU का device-gfx target इस्तेमाल करें – RX 6800M के लिए वह device-gfx1031 है:

python -m pip install --index-url https://repo.amd.com/rocm/whl-multi-arch/ \
    "torch[device-gfx1031]==2.12.0+rocm7.14.0" \
    "torchvision[device-gfx1031]==0.27.0+rocm7.14.0" \
    "torchaudio==2.11.0+rocm7.14.0"

कुछ notes:

अगर pip dependencies resolve करने में शिकायत करे, तो सुनिश्चित करें कि venv active है और आप --index-url https://repo.amd.com/rocm/whl-multi-arch/ pass कर रहे हैं – वही repository है जहाँ AMD ROCm builds publish करता है।

Step 4 – verify करें कि GPU detect हो रहा है

यह one-liner चलाकर confirm करें कि PyTorch आपका AMD GPU देख रहा है:

python -c "import torch; print(torch.cuda.is_available())"

अगर PyTorch और ROCm सही install हैं और आपका AMD GPU detect हो रहा है तो यह True print करता है। (हाँ, API में cuda ही है – PyTorch ROCm के लिए भी वही interface रखता है ताकि code portable रहे।)

आप आगे device का नाम भी check कर सकते हैं:

python -c "import torch; print(torch.cuda.get_device_name(0))"

RX 6800M पर यह कुछ ऐसा report करता है AMD Radeon RX 6800M

आगे क्या?

ROCm पर PyTorch चलने के साथ आप कर सकते हैं:

ज़्यादा detail चाहिए तो official guide देखें: Install PyTorch for ROCm – AMD AI ecosystem docs