torch.compile en profondeur

PyTorch 2.0 introduit torch.compile(), qui transforme un modèle eager en un graphe optimisé via TorchDynamo (capture) et TorchInductor (génération de code). Un seul appel suffit pour obtenir des accélérations de 30 à 200 % sans modifier le code du modèle.

Utilisation

python
import torch

model = torch.nn.TransformerEncoderLayer(
    d_model=512, nhead=8, batch_first=True
)
model = model.cuda()

# Compilation : une seule ligne
compiled = torch.compile(model, mode='reduce-overhead')

# Modes disponibles :
# 'default'         : bon compromis
# 'reduce-overhead' : minimise le surcoût CPU
# 'max-autotune'    : essaie plus de configs CUDA

x = torch.randn(32, 128, 512, device='cuda')
out = compiled(x)  # premier appel : compilation
out = compiled(x)  # appels suivants : rapides
print(out.shape)  # torch.Size([32, 128, 512])

Sources