Vue d'ensemble

PyTorch 1.13, publié le 29 octobre 2022, stabilise BetterTransformer et intègre functorch dans le noyau PyTorch.

Fonctionnalités principales

BetterTransformer

BetterTransformer accélère l'inférence des modèles Transformer avec des noyaux fusionnés et un support natif des séquences imbriquées.

python
import torch
from torch.nn import TransformerEncoderLayer

# BetterTransformer pour inférence rapide
layer = TransformerEncoderLayer(
    d_model=512, nhead=8, batch_first=True
)
fast_layer = torch.nn.utils.parametrize.transfer_parametrizations_and_params(
    layer, layer  # BetterTransformer activé automatiquement
)
x = torch.randn(2, 10, 512)
output = layer(x)
print(output.shape)  # (2, 10, 512)

functorch stable

functorch est désormais intégré dans PyTorch core, rendant vmap, grad et jacrev accessibles directement.

python
import torch
from torch.func import vmap, grad

# functorch dans torch.func
def loss_fn(x):
    return (x ** 2).sum()

# Gradient fonctionnel
x = torch.tensor([1.0, 2.0, 3.0])
grads = grad(loss_fn)(x)
print(grads)  # tensor([2., 4., 6.])

# Vectorisation
batch = torch.randn(5, 3)
results = vmap(grad(loss_fn))(batch)
print(results.shape)  # (5, 3)

Sources