PyTorch 搭配 DirectML 可在支援 DirectX 12 的 GPU 上進行訓練與推論。 PyTorch 搭配 DirectML 目前處於公開預覽版,並可於原生 Windows 上運作,從 Windows 10 1709 版本開始。
檢查你使用的 Windows 版本
要檢查你的 Windows 版本和組號,請選擇 Windows 標誌鍵 + R,輸入 winver,然後選擇確定。 如果你的版本較舊,請更新到 Windows 10 版本 1709 或更新。
檢查 GPU 驅動程式更新
透過 Windows Update 或硬體製造商的官方網站安裝顯示卡最新驅動程式。
設定 Python
安裝 Python 環境。 如果你用 Miniconda,請下載並執行你架構的 Windows 安裝程式。
接著建立並啟動一個名為 pytorch-directml 的環境:
conda create --name pytorch-directml python=3.10
conda activate pytorch-directml
使用 DirectML 安裝 PyTorch
安裝 torch-directml 套件:
pip install torch-directml
確認安裝情況
啟動 Python,然後執行下列程式碼,在 DirectML 裝置上將兩個張量相加:
import torch
import torch_directml
dml = torch_directml.device()
tensor1 = torch.tensor([1]).to(dml)
tensor2 = torch.tensor([2]).to(dml)
dml_algebra = tensor1 + tensor2
print(dml_algebra.item())
預期產出:
3
取樣與回饋
請參考 DirectML PyTorch 範例 作為範例。 若要回報套件問題或請求功能,請使用 DirectML 問題追蹤器。