Overview

PyTorch 1.13, released on October 29, 2022, stabilizes BetterTransformer and integrates functorch into PyTorch core.

Main Features

BetterTransformer

BetterTransformer accelerates Transformer inference with fused kernels and native nested sequence support.

python
import torch
from torch.nn import TransformerEncoderLayer

# BetterTransformer for fast inference
layer = TransformerEncoderLayer(
    d_model=512, nhead=8, batch_first=True
)
fast_layer = torch.nn.utils.parametrize.transfer_parametrizations_and_params(
    layer, layer  # BetterTransformer auto-enabled
)
x = torch.randn(2, 10, 512)
output = layer(x)
print(output.shape)  # (2, 10, 512)

Stable functorch

functorch is now integrated into PyTorch core, making vmap, grad, and jacrev directly accessible.

python
import torch
from torch.func import vmap, grad

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

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

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

Sources