Habilitar o PyTorch com DirectML no Windows

O PyTorch com DirectML permite treinamento e inferência em GPUs compatíveis com DirectX 12. O PyTorch com DirectML está em versão prévia pública e funciona em Windows nativos começando com Windows 10 versão 1709.

Verifique sua versão do Windows

Para verificar a versão do Windows e o número da compilação, pressione tecla do logotipo do Windows + R, digite winver e selecione OK. Atualize para Windows 10 versão 1709 ou posterior se o build for mais antigo.

Verificar se há atualizações de driver de GPU

Instale o driver mais recente disponível para sua GPU por meio de Windows Update ou do site do fabricante de hardware.

Configurar Python

Instale um ambiente de Python. Se você usar o Miniconda, baixe e execute o instalador de Windows para sua arquitetura.

Em seguida, crie e ative um ambiente chamado pytorch-directml:

conda create --name pytorch-directml python=3.10
conda activate pytorch-directml

Instalar o PyTorch com DirectML

Instalar o pacote torch-directml:

pip install torch-directml

Verificar a instalação

Inicie Python e execute o seguinte código para adicionar dois tensores no dispositivo 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())

Resultado esperado:

3

Exemplos e comentários

Confira os exemplos de PyTorch do DirectML para obter exemplos. Para relatar problemas de pacote ou recursos de solicitação, use o rastreador de problemas do DirectML.