Aller au contenu

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_CHECK sur 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");
ou, côté Python : 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_inline pour prototyper du CUDA C++, extension compilée pour distribuer.
  • torch.library.custom_op est ce qui rend un noyau utilisable par torch.compile sans graph break, avec autograd et mode meta.
  • torch.library.opcheck vé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 gradcheck en 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 :

\[\frac{\partial f}{\partial x} \approx \frac{f(x+\epsilon) - f(x-\epsilon)}{2\epsilon}\]

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 :

  1. attend la fin de toutes les opérations PyTorch en cours ;
  2. 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


Sources de ce chapitre