Vue d'ensemble
PyTorch 2.0, publié le 16 mars 2023, est une version majeure qui introduit torch.compile() pour l'optimisation automatique des graphes, offrant des gains de performance de 30 à 200% sans modification du code.
Fonctionnalités principales
torch.compile()
torch.compile() capture et optimise automatiquement le graphe de calcul PyTorch. Il analyse le code Python et génère du code optimisé pour le GPU, sans changer l'API existante.
import torch
model = torch.nn.Sequential(
torch.nn.Linear(784, 256),
torch.nn.ReLU(),
torch.nn.Linear(256, 10),
)
# Compilation automatique du modèle
model_compile = torch.compile(model)
# Utilisation identique, performances améliorées
x = torch.randn(32, 784)
y = model_compile(x) # 30-200% plus rapide
Optimisation des graphes
Le nouveau backend TorchDynamo capture le graphe de calcul au niveau Python et le passe à TorchInductor pour la génération de code optimisé (Triton pour GPU, C++ pour CPU).
import torch
@torch.compile(mode='reduce-overhead')
def train_step(model, x, y, optimizer, loss_fn):
pred = model(x)
loss = loss_fn(pred, y)
loss.backward()
optimizer.step()
optimizer.zero_grad()
return loss
# Les modes disponibles :
# 'default' - bon compromis
# 'reduce-overhead' - minimise la latence
# 'max-autotune' - optimise au maximum
Gains de performance
Les benchmarks montrent des accélérations significatives sur des modèles variés : 43% en moyenne sur les modèles HuggingFace, 46% sur TorchBench et 26% sur les modèles TIMM, le tout sans modification du code.
import torch
import time
model = torch.nn.Transformer(
d_model=512, nhead=8, num_encoder_layers=6
)
x = torch.randn(10, 32, 512)
# Sans compilation
start = time.time()
for _ in range(100):
model(x, x)
print(f'Sans compile : {time.time() - start:.2f}s')
# Avec compilation
model_c = torch.compile(model)
model_c(x, x) # warm-up
start = time.time()
for _ in range(100):
model_c(x, x)
print(f'Avec compile : {time.time() - start:.2f}s')
