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])
