Vue d'ensemble

PyTorch 1.11, publié le 11 février 2022, introduit TorchData pour le chargement de données et intègre functorch pour les transformations fonctionnelles.

Fonctionnalités principales

TorchData

TorchData fournit des DataPipes composables pour construire des pipelines de chargement de données flexibles et performants, remplaçant progressivement les anciens DataLoader.

python
from torchdata.datapipes.iter import IterableWrapper

# Pipeline de données composable
dp = IterableWrapper(range(100))
dp = dp.shuffle(buffer_size=20)
dp = dp.batch(batch_size=8)
dp = dp.map(lambda batch: [x * 2 for x in batch])

for batch in dp:
    print(batch)  # [liste de 8 éléments doublés]
    break

functorch

functorch apporte des transformations fonctionnelles comme vmap (vectorisation automatique) et grad (différentiation), inspirées de JAX.

python
import torch
from functorch import vmap, grad

# Vectorisation automatique avec vmap
def compute(x):
    return torch.sum(x ** 2)

batch = torch.randn(10, 3)
results = vmap(compute)(batch)  # appliqué à chaque ligne
print(results.shape)  # (10,)

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

Sources