PyTorch 搭配 DirectML 可在 Windows 子系統 Linux 版(WSL)中支援 DirectX 12 的 GPU 上進行訓練與推論。 PyTorch 搭配 DirectML 目前處於公開預覽階段,並可在 WSL 2 中運作。
檢查你使用的 Windows 版本
torch-directml WSL 2 中的套件需要 Windows 11,版本為 22000 或更新版本。 要檢查你的 Windows 版本和組號,請選擇 Windows 標誌鍵 + R,輸入 winver,然後選擇確定。
安裝 WSL 2
要安裝預設的 Linux 發行版與 WSL 2,請以管理員模式開啟 PowerShell 或 Windows 命令提示字元並執行:
wsl --install
系統提示時,請重新啟動您的電腦。 關於發行版選擇及其他安裝選項,請參見「使用 WSL 在 Windows 上安裝 Linux」。
檢查 GPU 驅動程式更新
請透過 Windows Update 或硬體製造商官網安裝最新 Windows 驅動程式。 Windows 驅動程式在 WSL 中啟用 GPU 加速;你不需要安裝獨立的 Linux 顯示驅動程式。
設定 Python
在你的 WSL 發行版中安裝 Python 環境。 例如,執行以下指令來安裝 Miniconda:
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
接著建立並啟用名為 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 問題追蹤器。