Documentation

compile=True et la voie rapide TileLang : notes d’ingénierie

Ces notes couvrent le comportement de compile=True et fast=True au-delà de ce que dit le README. Elles proviennent de mesures prises en travaillant sur #472, #576 et #718, sur une RTX 4070 Ti SUPER avec torch 2.11 et tilelang 0.1.14. Elles sont là pour que la personne suivante n’ait pas à les mesurer à nouveau.

Sélection du backend

Agent(..., backend="auto") et laya.load(..., backend="auto") optent pour la couche de classes de backend. La valeur par défaut reste eager. Un backend= explicite prime sur compile et fast ; l’omettre conserve le comportement existant des deux flags.

  • eager : le forward PyTorch standard, sur tout appareil pris en charge.
  • compile : torch.compile CUDA uniquement avec formes dynamiques et mode reduce-overhead, remplissage par buckets, cache inductor persistant et warmup à l’installation. Il réutilise la même portée de dimension indépendante que compile=True, qui garde son mode par défaut existant et le support CPU. Mets LAYA_COMPILE_WARMUP=0 pour différer le warmup du backend et LAYA_INDUCTOR_CACHE_DIR pour choisir son répertoire de cache (défaut ~/.cache/laya/inductor).
  • tilelang : un adaptateur autour de la voie rapide actuelle, utilisant le dtype bf16 ou fp16 de l’agent.
  • auto : TileLang sur CUDA avec un encodeur ModernBERT et un dtype pris en charge quand TileLang est installé, sinon compile sur CUDA ; eager sur les autres appareils.
  • onnx : laya.load(..., backend="onnx", onnx_path="model.onnx") renvoie l’ONNXAgent existant. Sans onnx_path, il utilise laya.onnx.

Un backend indisponible émet un RuntimeWarning nommant le backend résolu et retombe sur eager. Pour exiger un backend, utilise agent.set_backend("tilelang", strict=True). Le changement attend l’inférence active ; agent.backend rapporte le nom actif et agent.backend_object expose l’objet installé. agent.set_backend("compile", warmup=False) diffère la compilation jusqu’à l’inférence, donc les erreurs de compilation apparaissent alors sur la requête. agent.warmup() reste disponible. agent.deaccelerate() retire un backend installé via la couche de classes.

Les Router transmettent une sélection explicite via Router(agent_kwargs={"backend": "auto"}). Ils ne passent aucun argument de backend par défaut, préservant la compatibilité avec les constructeurs de type Agent existants. Les reprises CPU OOM à portée limitée détachent le backend et le restaurent quand le modèle revient sur son appareil d’origine.

compile=True matérialise le masque d’attention

Eager SDPA prend le masque d’attention (rows, 1, L, L) de ModernBERT comme une vue diffusée. Sous les formes dynamiques qu’utilise compile=True, inductor ne peut pas prouver que la dernière dimension est alignée. Il étend le masque à chaque tête et le remplit dans un vrai tampon de rows x heads x L x L. En bf16 avec 12 têtes, cela fait 0.8 GB à 32 lignes x 1024 jetons.

  • GPU avec de la marge. Le tampon coûte de la bande passante, des dizaines de ms par long lot.
  • GPU presque plein. L’allocateur de cache sature, et le même appel peut prendre des dizaines de secondes.

Si tu compiles avec de longs lots sur un GPU occupé, limite la taille de lot (predict_batch(..., batch_size=)) ou utilise fast=True. L’attention TileLang lit le tampon QKV empaqueté et masque par longueur de séquence, donc elle n’a pas ce tampon.

Démarrage à froid

  • Première compilation. Elle prend des dizaines de secondes par graphe. compile=True a besoin de deux graphes : un pour les lots et un pour une seule ligne, que torch spécialise. compile=True appelle maintenant agent.warmup() pendant le chargement. compile_warmup=False restaure la compilation paresseuse, et agent.warmup(shapes=...) reste disponible manuellement. Les chargements eager et TileLang ne s’échauffent pas automatiquement. Ces formes couvrent les requêtes courantes, pas toutes les gardes de forme possibles.
  • Échec du warmup. Le warmup automatique est au mieux : un échec émet un RuntimeWarning nommant l’erreur (y compris l’erreur du compilateur sous-jacente) et le chargement revient avec le wrapper torch.compile et les réglages de compilation intacts. Par exemple, Windows sans MSVC peut charger avec compile=True même si le warmup échoue. Les requêtes ultérieures utilisent toujours le modèle compilé et font apparaître les échecs de compilation ; Laya ne les bascule pas en exécution eager. Les appels explicites à agent.warmup() propagent aussi les échecs, y compris après un warmup automatique échoué. Un chargement réussi ne garantit donc pas que l’inférence compilée est prête.
  • Opt-in du cache Laya. laya.load(..., compile=True, compile_cache=True) définit le TORCHINDUCTOR_CACHE_DIR global au processus seulement quand il est absent, sur $XDG_CACHE_HOME/laya/torchinductor ou ~/.cache/laya/torchinductor quand XDG n’est pas défini ou n’est pas absolu. Un réglage existant, y compris celui posé par une compilation PyTorch antérieure, l’emporte. Le répertoire est créé au chargement ; les erreurs de système de fichiers se propagent. compile_cache=False (défaut), eager et les chargements TileLang laissent l’environnement tranquille. Cela ne déplace ni ne supprime les anciens caches. Les conteneurs ont toujours besoin d’un home/volume persistant. La compatibilité et l’invalidation du cache sont gérées par PyTorch ; un changement de GPU, torch, compilateur, modèle ou garde d’entrée peut exiger une recompilation.
  • Entre redémarrages. Le cache de graphes FX d’Inductor conserve les graphes compilés sous TORCHINDUCTOR_CACHE_DIR. La valeur par défaut est sous /tmp, qui ne survit ni à un redémarrage ni à un redémarrage de conteneur. Définis-le sur un répertoire persistant, ou un volume dans un conteneur, et un second processus charge les graphes au lieu de les compiler. Dans la mesure de #472, cela a fait passer le préchauffage d’environ 120 s à environ 50 s.

Graphes CUDA en option

agent = laya.load("convaiinnovations/laya", compile=True,
                  compile_cache=True, compile_mode="reduce-overhead")

compile_mode vaut "default" par défaut ; sur le chemin compilé actif, seuls "default" et "reduce-overhead" sont acceptés. Les chargements eager et TileLang ignorent les options de compilation. La compilation CPU fonctionne toujours, mais l’enregistrement de graphes CUDA ne s’applique que sur CUDA. Le mode CUDA nécessite l’API torch.compiler.cudagraph_mark_step_begin de PyTorch ; les builds plus anciens sans elle lèvent une erreur explicite.

Les graphes Dynamo dynamiques n’impliquent pas des graphes CUDA indépendants de la forme : de nouvelles formes concrètes peuvent nécessiter à nouveau warmup et enregistrement, sans nouveau graphe Dynamo. Les deux formes de warmup synthétiques par défaut ne pré-enregistrent pas toutes les formes de requête. Les formes répétées peuvent en bénéficier, mais les formes variables peuvent payer une latence supplémentaire et retenir des pools de graphes. PyTorch peut ignorer les graphes CUDA pour des opérations ou configurations non prises en charge ; définir ce mode n’est pas une garantie de capture.

Laya marque chaque forward CUDA compilé comme une nouvelle étape, sérialise ces forwards entre ses agents et clone les deux tenseurs de sortie hors du graphe compilé avant de libérer le verrou. Cela garde les sorties retenues valides à travers les replays, au prix de deux copies et d’une exécution de forward sérialisée. Le verrou ne coordonne pas les modèles compilés indépendants appartenant à l’application ; les appelants qui partagent des itérations de graphes CUDA ou utilisent des streams personnalisés doivent gérer leur propre coordination. Les caches disque réutilisent le code compilé, pas les enregistrements vivants de graphes CUDA ni leur mémoire d’appareil, entre processus.

Reproduis les timings de froid/redémarrage, la mémoire et les compteurs de cache avec benchmarks/bench_compile_defaults.py ; voir les mesures enregistrées.

AOTInductor : pas encore

Livrer un artefact précompilé par checkpoint et par architecture GPU (torch._inductor.aoti_compile_and_package) supprimerait complètement la compilation. Sur torch 2.11, cela s’arrête au packaging :

  • L’export fonctionne. torch.export de DecisionModel réussit, en environ 5 s, avec des lignes, des marqueurs et des jetons dynamiques. Les jetons doivent être déclarés comme multiple de 16 (16 * Dim(...)) ; une plage simple fait échouer la propre garde d’alignement L % 8 de l’exportateur. C’est le même alignement de masque que ci-dessus.
  • Le packaging échoue. La façon dont il échoue dépend de la façon dont le programme a été exporté :
    • Sous autocast, le programme porte des assertions de dtype sur lesquelles AOTI trébuche hors autocast : Tensor dtype mismatch! Expected: torch.bfloat16, Got: torch.float32.
    • Depuis une copie bf16 sans autocast, le traçage échoue à l’intérieur de la passe avant : mat1 and mat2 must have the same dtype. DecisionModel.forward convertit l’état regroupé et les caractéristiques de confiance en fp32 avant la tête d’action, et autocast normalement réconcilie cela.

La voie de l’artefact a donc besoin d’une tête d’action à dtype explicite : soit elle convertit son entrée vers le dtype de la tête, soit elle exécute la tête en fp32.

Portabilité de TileLang : les kernels sont CUDA uniquement

tilelang enregistre des cibles pour CUDA, HIP, Metal, WebGPU et un backend C. Sans matériel AMD ni Apple, la question à laquelle on pouvait répondre était de savoir si laya/tl_kernels.py compile pour le CPU tout court. Sondé avec tilelang.compile(kernel.prim_func, target=...) sous Linux x86-64 :

cible résultat
"cpu" Rejeté d’emblée : Target cpu is not supported. Le backend CPU de tilelang est "c".
"llvm" Cannot find global function target.build.llvm. La roue n’inclut pas de backend LLVM.
"c" Compile vers C et s’exécute sur des tenseurs CPU, mais seulement pour un sous-ensemble du langage.

Tous les kernels Laya échouent avec "c", pour l’une de trois raisons :

kernel échec avec target="c" construction
gemm_kernel, gemm_geglu_kernel CPU fill only supports local and global buffers, but got dst scope local.fragment accumulateur T.alloc_fragment
add_ln_kernel CPU reduce only supports local src and local/local.var dst buffers T.reduce_sum / T.reduce_max sur des fragments
rope_kernel Cannot convert type bfloat16 to C type tenseurs bf16
attn_kernel échoue à T.alloc_fragment fragments

Le backend C accepte bien :

  • les boucles élément par élément en fp32 (T.Parallel) ;
  • T.Pipelined, qui compile en une boucle simple ;
  • T.gemm avec un accumulateur T.alloc_local, qui compile en une boucle scalaire triple.

Une version CPU serait donc un second jeu de kernels, pas un drapeau de cible. Son GEMM serait une boucle scalaire non bloquée, et elle ne rivaliserait pas avec la voie MKL/oneDNN que la passe avant standard utilise déjà sur CPU. Les trois mêmes constructions sont celles à vérifier en premier sur HIP et Metal : les fragments, T.gemm avec GemmWarpPolicy, et le support bf16/fp16.

Pour reproduire la première ligne du second tableau : tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c").