在 Windows 上啟用 PyTorch 搭配 DirectML

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 問題追蹤器。