在 WSL 上使用 DirectML 啟用 PyTorch

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