Aller au contenu

5 · Helion et torch.compile

Le haut du gradient. Ce que le compilateur fait tout seul, et ce qu'un DSL de très haut niveau peut faire de plus.


5.1 torch.compile, en une page

modele = torch.compile(modele)

Ce que cela déclenche :

Code Python
    │
    ├─ TorchDynamo   : capture le graphe (bytecode Python → FX graph)
    │                   avec "graph breaks" sur ce qu'il ne sait pas tracer
    ├─ AOTAutograd   : dérive le graphe backward
    ├─ TorchInductor : optimise et génère du code
    │                   ├─ fusion des opérations élémentaires
    │                   ├─ choix des noyaux de GEMM (cuBLAS / CUTLASS / Triton)
    │                   ├─ planification mémoire
    │                   └─ génération de TRITON
    └─ CUDA Graphs   : si mode="reduce-overhead"

Le point crucial : Inductor génère du Triton. Quand vous utilisez torch.compile, vous exécutez des noyaux Triton générés automatiquement.

Pour les lire :

import os
os.environ["TORCH_COMPILE_DEBUG"] = "1"
os.environ["TORCH_LOGS"] = "output_code"
# les noyaux Triton générés sont écrits dans /tmp/torchinductor_<user>/

C'est un excellent exercice pédagogique : comparer votre noyau manuel à celui qu'Inductor produit.

Les modes

Mode Ce qu'il fait Coût de compilation
default fusion + sélection de noyaux quelques secondes
reduce-overhead + CUDA Graphs idem, mais formes figées
max-autotune + réglage automatique des GEMM et des convolutions minutes

Ce que ça gagne

Les ordres de grandeur constatés :

Charge Gain typique
Modèle avec beaucoup d'opérations élémentaires 1,5× à 2,5×
Modèle dominé par de grosses GEMM 1,05× à 1,2×
Inférence à petit lot 1,3× à 2× (surtout grâce aux CUDA Graphs)
Entraînement 1,2× à 1,8×

Le gain vient principalement de la fusion élémentaire : là où PyTorch en mode impératif lance un noyau par opération, Inductor en lance un seul.

Les limites

Les cinq ennuis de torch.compile

  1. Les graph breaks — tout ce que Dynamo ne sait pas tracer (I/O, print, structures de données Python complexes, conditions dépendant des données) coupe le graphe et réduit la fusion.

    torch._dynamo.explain(modele)(entree)   # liste les breaks
    
  2. Les recompilations — un changement de forme d'entrée déclenche une recompilation. dynamic=True aide, torch._dynamo.config.cache_size_limit aussi.

  3. Le temps de compilation — max-autotune peut prendre plusieurs minutes sur un gros modèle.
  4. reduce-overhead fige les formes et les pointeurs (CUDA Graphs) : incompatible avec des lots de taille variable sans précautions.
  5. Inductor ne fusionne pas tout — les motifs complexes (attention, fusions à travers une réduction) restent à votre charge.

5.2 Helion

Helion est un DSL Python développé par Meta, présenté à la PyTorch Conference 2025. Sa position : entre PyTorch et Triton.

Helion can be viewed either as PyTorch with tiles or as a higher-level Triton.

Il compile vers Triton.

La syntaxe

import helion
import helion.language as hl
import torch

@helion.kernel
def matmul(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
    m, k = x.size()
    k2, n = y.size()
    sortie = torch.empty([m, n], dtype=x.dtype, device=x.device)

    for tile_m, tile_n in hl.tile([m, n]):          # ← boucle sur les tuiles
        acc = hl.zeros([tile_m, tile_n], dtype=torch.float32)
        for tile_k in hl.tile(k):
            acc = torch.addmm(acc, x[tile_m, tile_k], y[tile_k, tile_n])
        sortie[tile_m, tile_n] = acc

    return sortie

Trois observations :

  1. Les tailles de tuile ne sont pas spécifiées. hl.tile([m, n]) déclare qu'on itère par tuiles ; l'autotuner choisit les dimensions.
  2. Les opérations sont du PyTorch (torch.addmm, indexation) appliquées à des tuiles.
  3. L'espace de recherche est implicite — c'est l'argument principal de Helion : dans Triton il faut énumérer les configurations à la main, ici elles sont dérivées.

Les performances annoncées

Comparaison Gain moyen
Helion contre torch.compile 1,05×
Helion contre noyaux Triton écrits à la main 1,44×
Cas extrêmes (int4_gemm, jsd) 4,5× et 4,4×

Source : PyTorch blog.

Comment lire ces chiffres

« 1,44× contre du Triton écrit à la main » ne signifie pas que Helion génère du meilleur code que Triton — il génère du Triton. Cela signifie que son autotuner explore un espace plus large que ce qu'un humain énumère manuellement.

C'est une différence importante : le gain vient du réglage, pas de la génération. Un noyau Triton exhaustivement réglé à la main atteindrait la même performance, pour beaucoup plus d'effort.

Les cas à 4,5× correspondent probablement à des noyaux de référence peu optimisés. Ce sont des chiffres du fournisseur, non reproduits ici.

La portabilité

Helion cible NVIDIA et AMD. AMD documente d'ailleurs son usage sur ses propres GPU dans son AI Developer Hub, ce qui est un signe de portabilité réelle plutôt qu'annoncée.


5.3 Le tableau de décision

Votre modèle est lent
    │
    ├─ Avez-vous essayé torch.compile ?
    │      NON → essayez. 30 secondes d'effort.
    │      OUI ↓
    │
    ├─ Y a-t-il des graph breaks ?
    │      OUI → corrigez-les d'abord (torch._dynamo.explain)
    │      NON ↓
    │
    ├─ Le profileur montre-t-il un noyau dominant qu'Inductor a mal généré ?
    │      NON → le problème est ailleurs (données, communication, mémoire)
    │      OUI ↓
    │
    ├─ Est-ce une opération standard (GEMM, attention, convolution) ?
    │      OUI → utilisez la bibliothèque dédiée
    │      NON ↓
    │
    ├─ Helion : écrivez-le en 20 lignes, laissez l'autotuner travailler
    │      Assez rapide ? → terminé
    │      NON ↓
    │
    ├─ Triton : contrôle des tuiles et des étages de pipeline
    │      Assez rapide ? → terminé
    │      NON ↓
    │
    └─ Gluon / CuTe DSL / CUDA C++

5.4 Un mot sur JAX et Pallas

L'écosystème JAX a son propre chemin, structurellement analogue :

PyTorch JAX
torch.compile / Inductor jax.jit / XLA
Triton Pallas
Noyau CUDA custom jax.ffi / CuTe DSL

Pallas est le DSL de noyaux de JAX. Il a deux back-ends : Triton (GPU) et Mosaic (TPU), ce qui en fait le seul DSL de noyaux à cibler sérieusement les deux.

import jax
from jax.experimental import pallas as pl

def add_kernel(x_ref, y_ref, o_ref):
    o_ref[...] = x_ref[...] + y_ref[...]

@jax.jit
def add(x, y):
    return pl.pallas_call(
        add_kernel,
        out_shape=jax.ShapeDtypeStruct(x.shape, x.dtype),
        grid=(x.shape[0] // 128,),
        in_specs=[pl.BlockSpec((128,), lambda i: (i,)),
                  pl.BlockSpec((128,), lambda i: (i,))],
        out_specs=pl.BlockSpec((128,), lambda i: (i,)),
    )(x, y)

Le modèle est le même : des tuiles, des BlockSpec qui décrivent la correspondance grille → tuile, et le compilateur gère le reste.

JAX documente également l'appel de noyaux CuTe DSL depuis JAX, ce qui donne accès au niveau bas sans quitter l'écosystème.


Résumé du chapitre

À retenir

  • torch.compile = Dynamo (capture) + AOTAutograd (backward) + Inductor (génère du Triton) + éventuellement CUDA Graphs.
  • Gain typique : 1,5-2,5× sur les modèles riches en opérations élémentaires, 1,05-1,2× sur ceux dominés par les GEMM.
  • Vérifiez les graph breaks avant tout : ils annulent la fusion.
  • Helion = PyTorch avec des tuiles, compile vers Triton, autotuning à espace implicite. 1,05× contre torch.compile, 1,44× contre du Triton manuel — le gain vient du réglage, pas de la génération.
  • Côté JAX : jax.jit/XLA en haut, Pallas au niveau de Triton (avec un back-end TPU), CuTe DSL en dessous.
  • Descendez le gradient seulement quand le profileur désigne un noyau précis.

Vérifiez que vous avez compris

Pourquoi torch.compile aide-t-il peu sur un modèle dominé par de grosses GEMM ?

Parce que les grosses GEMM sont déjà exécutées par cuBLAS, qui est proche du pic. Inductor ne peut ni les fusionner utilement (elles sont limitées par le calcul, pas par la mémoire) ni les remplacer par du Triton plus rapide.

Le gain d'Inductor vient de l'élimination des allers-retours en HBM sur les opérations élémentaires. Si celles-ci représentent 5 % du temps, le gain maximal est de 5 %.

C'est une application directe du roofline : on ne gagne que là où on est limité par la mémoire.

Helion « bat les noyaux Triton écrits à la main de 1,44× ». Comment est-ce possible s'il génère du Triton ?

Parce que la comparaison ne porte pas sur la qualité du code généré mais sur la qualité du réglage.

En Triton, le développeur énumère manuellement quelques configurations (triton.Config). En Helion, l'espace de recherche est dérivé de la structure du noyau et exploré automatiquement — il couvre donc des points que l'humain n'a pas pensé à tester.

Corollaire honnête : un noyau Triton réglé exhaustivement atteindrait la même performance. Helion vend du temps d'ingénieur, pas une meilleure compilation.

Vous voyez Recompiling function forward à chaque itération. Que se passe-t-il ?

Dynamo garde le code compilé sous des conditions (formes, types, valeurs de certaines variables Python). Si une condition change, il recompile.

Causes fréquentes :

  • formes d'entrée variables → torch.compile(modele, dynamic=True) ;
  • une valeur Python utilisée dans une branche change à chaque itération (un compteur, un pas d'apprentissage lu depuis un objet) ;
  • dépassement de cache_size_limit (8 par défaut), qui provoque un abandon complet et un retour au mode impératif.

Diagnostic :

import torch._logging
torch._logging.set_logs(recompiles=True)

Chapitre suivant : 6 · Mojo


Sources de ce chapitre