Habilitación de PyTorch con DirectML en Windows

PyTorch con DirectML permite el entrenamiento y la inferencia en GPU compatibles con DirectX 12. PyTorch con DirectML está en versión preliminar pública y funciona en Windows nativo a partir de Windows 10 versión 1709.

Comprobación de la versión de Windows

Para comprobar la versión de Windows y el número de compilación, seleccione tecla del logotipo de Windows + R, escriba winver y seleccione Aceptar. Actualice a Windows 10 versión 1709 o posterior si la compilación es anterior.

Buscar actualizaciones de controladores de GPU

Instale el controlador más reciente disponible para la GPU a través de Windows Update o el sitio web del fabricante de hardware.

Configuración de Python

Instale un entorno de Python. Si usa Miniconda, descargue y ejecute el instalador de Windows para la arquitectura.

A continuación, cree y active un entorno denominado pytorch-directml:

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

Instalación de PyTorch con DirectML

Instala el paquete torch-directml:

pip install torch-directml

Comprobación de la instalación

Inicie Python y ejecute el código siguiente para agregar dos tensores en el 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

Ejemplos y comentarios

Consulte los ejemplos de PyTorch de DirectML para obtener ejemplos. Para notificar problemas de paquetes o solicitar características, use el seguimiento de problemas de DirectML.