3 · Intégrer un noyau dans PyTorch¶
Quatre méthodes, du prototype d'une minute à l'extension distribuable, avec la gestion des gradients et l'enregistrement d'opérateurs personnalisés.
3.1 Le tableau de décision¶
| Méthode | Compilation | Usage |
|---|---|---|
| Triton | JIT, transparente | par défaut pour tout ce qui est fusion et réduction |
load_inline |
JIT au premier appel | prototypage de CUDA C++ |
| Extension compilée | à l'installation | production, distribution |
torch.library |
s'ajoute aux précédentes | intégration propre : autograd, torch.compile, meta |
3.2 Triton — le chemin par défaut¶
import torch
import triton
import triton.language as tl
@triton.jit
def gelu_kernel(x_ptr, y_ptr, n, BLOC: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOC + tl.arange(0, BLOC)
m = offs < n
x = tl.load(x_ptr + offs, mask=m)
# GELU approximée par tanh
y = 0.5 * x * (1.0 + tl.math.tanh(
0.7978845608 * (x + 0.044715 * x * x * x)))
tl.store(y_ptr + offs, y, mask=m)
def gelu(x: torch.Tensor) -> torch.Tensor:
y = torch.empty_like(x)
n = x.numel()
grille = lambda meta: (triton.cdiv(n, meta["BLOC"]),)
gelu_kernel[grille](x, y, n, BLOC=1024)
return y
Rien à compiler, rien à configurer. Le noyau est compilé au premier appel et mis en cache.
C'est le chemin par défaut, sauf si vous avez besoin de quelque chose que Triton ne sait pas faire (synchronisation inter-blocs, contrôle des layouts).
3.3 load_inline — prototyper du CUDA C++¶
import torch
from torch.utils.cpp_extension import load_inline
source_cuda = r"""
#include <torch/extension.h>
__global__ void gelu_kernel(const float* x, float* y, int n) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) {
float v = x[i];
y[i] = 0.5f * v * (1.0f + tanhf(0.7978845608f *
(v + 0.044715f * v * v * v)));
}
}
torch::Tensor gelu_cuda(torch::Tensor x) {
TORCH_CHECK(x.is_cuda(), "x doit être sur GPU");
TORCH_CHECK(x.is_contiguous(), "x doit être contigu");
TORCH_CHECK(x.scalar_type() == torch::kFloat32, "float32 attendu");
auto y = torch::empty_like(x);
int n = x.numel();
int threads = 256;
int blocs = (n + threads - 1) / threads;
gelu_kernel<<<blocs, threads>>>(
x.data_ptr<float>(), y.data_ptr<float>(), n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return y;
}
"""
declaration = "torch::Tensor gelu_cuda(torch::Tensor x);"
module = load_inline(
name="gelu_ext",
cpp_sources=declaration,
cuda_sources=source_cuda,
functions=["gelu_cuda"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
verbose=True,
)
x = torch.randn(1_000_000, device="cuda")
y = module.gelu_cuda(x)
torch.testing.assert_close(y, torch.nn.functional.gelu(x, approximate="tanh"),
rtol=1e-4, atol=1e-4)
Points d'attention :
- la compilation prend 10 à 60 secondes au premier appel, puis est mise en
cache (
~/.cache/torch_extensions/) ; C10_CUDA_KERNEL_LAUNCH_CHECK()vérifie l'erreur de lancement — ne l'omettez pas ;- les
TORCH_CHECKsur le périphérique, la contiguïté et le type sont obligatoires : sans eux, un tenseur CPU passé par erreur produit un segfault au lieu d'un message clair.
3.4 L'extension compilée¶
Pour distribuer.
setup.py :
from setuptools import setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
setup(
name="mes_noyaux",
ext_modules=[
CUDAExtension(
name="mes_noyaux._C",
sources=["src/liaisons.cpp", "src/noyaux.cu"],
extra_compile_args={
"cxx": ["-O3"],
"nvcc": ["-O3", "-lineinfo",
"-gencode", "arch=compute_80,code=sm_80",
"-gencode", "arch=compute_90,code=sm_90"],
},
)
],
cmdclass={"build_ext": BuildExtension},
)
src/liaisons.cpp :
#include <torch/extension.h>
torch::Tensor gelu_cuda(torch::Tensor x);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("gelu", &gelu_cuda, "GELU (CUDA)");
}
pip install -e .
3.5 torch.library — l'intégration propre¶
C'est ce qui distingue un noyau bricolé d'un opérateur de première classe.
import torch
# 1. Déclarer l'opérateur
@torch.library.custom_op("mes_noyaux::gelu", mutates_args=())
def gelu(x: torch.Tensor) -> torch.Tensor:
return module.gelu_cuda(x)
# 2. Fonction "meta" : permet à torch.compile de raisonner sans exécuter
@gelu.register_fake
def _(x):
return torch.empty_like(x)
# 3. Gradient
def _backward(ctx, grad):
(x,) = ctx.saved_tensors
return module.gelu_backward_cuda(grad, x)
def _setup_context(ctx, inputs, output):
(x,) = inputs
ctx.save_for_backward(x)
torch.library.register_autograd(
"mes_noyaux::gelu", _backward, setup_context=_setup_context)
Ce que cela apporte :
Avec torch.library |
Sans |
|---|---|
torch.compile sait le tracer |
graph break |
| L'autograd fonctionne | il faut une autograd.Function manuelle |
Le mode meta fonctionne (formes sans exécution) |
non |
torch.export et les backends fonctionnent |
non |
Testable avec opcheck |
non |
# Vérification automatique de conformité
torch.library.opcheck(torch.ops.mes_noyaux.gelu, (x,))
Le point crucial : les graph breaks
Un noyau custom non enregistré coupe le graphe de torch.compile. Inductor
ne peut plus fusionner à travers lui, et vous perdez souvent plus que ce que
votre noyau apporte.
L'enregistrement via torch.library élimine ce problème.
3.6 L'autograd manuel¶
Si vous n'utilisez pas torch.library :
class GeluFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return module.gelu_cuda(x)
@staticmethod
def backward(ctx, grad_sortie):
(x,) = ctx.saved_tensors
return module.gelu_backward_cuda(grad_sortie.contiguous(), x)
gelu = GeluFunction.apply
Vérifier le gradient — étape obligatoire :
from torch.autograd import gradcheck
x = torch.randn(64, dtype=torch.double, device="cuda", requires_grad=True)
assert gradcheck(gelu, (x,), eps=1e-6, atol=1e-4)
gradcheck compare le gradient analytique aux différences finies. Il exige du
float64 : en float32, le bruit numérique des différences finies noie le
signal.
3.7 Les cinq pièges¶
Ce qui casse en pratique
1. Les tenseurs non contigus. Une vue transposée n'a pas la disposition que suppose votre noyau.
TORCH_CHECK(x.is_contiguous(), "x doit être contigu");
x = x.contiguous().
2. Le mauvais flux CUDA. PyTorch utilise son propre flux ; lancer sur le flux nul crée une synchronisation implicite.
auto flux = at::cuda::getCurrentCUDAStream();
mon_noyau<<<g, b, 0, flux>>>(...);
3. Le mauvais périphérique. En multi-GPU, il faut activer le bon.
const at::cuda::OptionalCUDAGuard garde(device_of(x));
4. Les erreurs silencieuses. Toujours
C10_CUDA_KERNEL_LAUNCH_CHECK() après le lancement.
5. Le mauvais type. Utilisez AT_DISPATCH_FLOATING_TYPES_AND2 pour
gérer float, half et bfloat16 :
AT_DISPATCH_FLOATING_TYPES_AND2(
at::ScalarType::Half, at::ScalarType::BFloat16,
x.scalar_type(), "mon_noyau", [&] {
mon_noyau<scalar_t><<<g, b, 0, flux>>>(
x.data_ptr<scalar_t>(), y.data_ptr<scalar_t>(), n);
});
3.8 Mesurer correctement¶
import torch
def mesurer(f, *args, rep=100, echauffement=20):
for _ in range(echauffement):
f(*args)
torch.cuda.synchronize()
debut = torch.cuda.Event(enable_timing=True)
fin = torch.cuda.Event(enable_timing=True)
debut.record()
for _ in range(rep):
f(*args)
fin.record()
torch.cuda.synchronize()
return debut.elapsed_time(fin) / rep # millisecondes
x = torch.randn(1 << 24, device="cuda")
t_mien = mesurer(gelu, x)
t_torch = mesurer(lambda t: torch.nn.functional.gelu(t, approximate="tanh"), x)
octets = 2 * x.numel() * x.element_size() # lecture + écriture
print(f"Mien : {t_mien*1000:7.1f} µs {octets/t_mien/1e9:6.1f} Go/s")
print(f"PyTorch: {t_torch*1000:7.1f} µs {octets/t_torch/1e9:6.1f} Go/s")
print(f"Ratio : {t_torch/t_mien:.2f}×")
Il existe aussi triton.testing.do_bench, qui gère l'échauffement, le vidage du
cache L2 et les quantiles :
import triton
ms = triton.testing.do_bench(lambda: gelu(x), warmup=25, rep=100)
La métrique à afficher
Pour un noyau limité par la mémoire, affichez les Go/s et comparez à la bande passante crête de votre carte. « 2 900 sur 3 350 possibles, soit 87 % » est une information ; « 42 µs » n'en est pas une.
Résumé du chapitre¶
À retenir
- Triton par défaut.
load_inlinepour prototyper du CUDA C++, extension compilée pour distribuer. torch.library.custom_opest ce qui rend un noyau utilisable partorch.compilesans graph break, avec autograd et modemeta.torch.library.opcheckvérifie automatiquement la conformité.- Cinq pièges : contiguïté, flux CUDA de PyTorch, périphérique courant, vérification d'erreur de lancement, dispatch de type.
- Vérifiez le gradient avec
gradchecken float64. - Mesurez avec des événements CUDA après échauffement, et affichez la fraction de la bande passante crête.
Vérifiez que vous avez compris¶
Pourquoi gradcheck exige-t-il du float64 ?
Parce qu'il compare le gradient analytique à une différence finie :
Avec \(\epsilon = 10^{-6}\) et une précision float32 de \(\sim10^{-7}\), la
soustraction \(f(x+\epsilon) - f(x-\epsilon)\) perd presque tous ses chiffres
significatifs : c'est de l'annulation catastrophique.
En float64 (\(\sim10^{-16}\)), il reste largement assez de précision pour que
la comparaison ait un sens.
Conséquence pratique : votre noyau doit supporter le double pour être
testable, même s'il ne sert qu'en float32 ou bfloat16 en production.
Votre noyau custom est 1,5× plus rapide isolément, mais le modèle complet est plus lent après intégration. Pourquoi ?
Presque certainement un graph break dans torch.compile.
Un opérateur non enregistré via torch.library interrompt le graphe.
Inductor ne peut plus fusionner les opérations qui l'entourent, et vous
perdez plusieurs fusions élémentaires — chacune valant un aller-retour en
HBM.
Si votre noyau économise 20 µs mais casse trois fusions valant 15 µs chacune, le bilan est négatif de 25 µs.
Diagnostic :
torch._dynamo.explain(modele)(entree) # liste les breaks
Remède : enregistrer l'opérateur avec custom_op et register_fake.
Pourquoi lancer sur at::cuda::getCurrentCUDAStream() plutôt que sur le flux par défaut ?
Parce que PyTorch exécute tout sur ses propres flux, et que le flux nul a une sémantique de synchronisation implicite : une opération sur le flux nul attend que tous les autres flux soient vides, puis les bloque.
Concrètement, votre noyau lancé sur le flux nul :
- attend la fin de toutes les opérations PyTorch en cours ;
- bloque les opérations suivantes jusqu'à sa propre fin.
Vous perdez tout recouvrement, et vous introduisez une barrière invisible au milieu du modèle. Sur un modèle avec beaucoup d'opérations concurrentes, cela peut coûter plus que ce que votre noyau apporte.
Chapitre suivant : 4 · Tester et déboguer