Skip to content

ubatch et cudagraph : comment les graphes à forme fixe accélèrent decode/prefill

源码版本v0.25.1

Responsabilités

cudagraph est le mécanisme CUDA « enregistrer une fois, rejouer à l'infini » : on grave une séquence de kernels en graphe, et chaque replay économise le coût de kernel launch. En phase de decode, chaque token ne calcule qu'un ou deux nouveaux tokens, les kernels sont petits et nombreux — le terrain idéal pour cudagraph. v1 utilise CUDAGraphWrapper(cuda_graph.py:145-233) pour envelopper le forward du modèle. À la première rencontre d'un BatchDescriptor, __call__ enregistre le graphe(cuda_graph.py:265-344) ; ensuite, pour la même forme, on appelle directement entry.cudagraph.replay()(cuda_graph.py:357-361).

ubatch (micro-batch) est le mécanisme de v1 : quand DP>1, on découpe le grand batch d'un step en N petits batchs à forme fixe, chacun enregistré dans son propre graphe ; les différents DP ranks peuvent alors s'exécuter en parallèle et faire leurs communications TP/EP all2all simultanément, en chevauchant communication et calcul. Cette logique est implémentée dans UBatchWrapper(gpu_ubatch_wrapper.py:113-211) et UBatchContext(ubatching.py:20-148). Ensemble, les deux forment le DBO (Distributed Batch Overlap) : découpage en micro-batchs + capture cudagraph + coordination par sémaphores d'événements multithreads.

Le mode d'exécution cudagraph est décidé par CUDAGraphMode(compilation.py:53-87) : NONE (pas d'enregistrement), PIECEWISE (par morceaux, attention hors graphe), FULL (graphe complet), FULL_DECODE_ONLY (FULL seulement en decode, prefill en piecewise), etc. À la fin de GPUModelRunner.load_model, selon la configuration, on enroule self.model dans le wrapper correspondant(gpu_model_runner.py:5350-5379).

Motivation de conception

  • cudagraph piloté par la forme : CUDAGraphWrapper.__call__ utilise BatchDescriptor comme clé(cuda_graph.py:257-263), chaque forme n'est enregistrée qu'une fois, puis on rejoue ; CUDAGraphDispatcher centralise dans le runner « quel mode pour quelle forme »(gpu_model_runner.py:844-845) ; _determine_batch_execution_and_padding pad num_tokens jusqu'à la forme enregistrable la plus proche(gpu_model_runner.py:3879-3894).
  • PIECEWISE saute l'attention : dans un graphe FULL, le kernel d'attention a une forme trop malléable (query/key de longueur variable), l'enregistrer ne sert à rien au replay ; v1 propose donc le mode piecewise, qui découpe le modèle en segments, l'attention en eager, le reste dans le graphe, via BreakableCUDAGraphWrapper(breakable_cudagraph.py:246).
  • ubatch à découpe uniforme : maybe_create_ubatch_slices découpe au point num_tokens_padded // num_ubatches(ubatch_utils.py:63-114), chaque ubatch a la même forme, de sorte que le même cudagraph peut être rejoué plusieurs fois et que les DP ranks s'alignent.
  • Coordination multi-thread + événements CUDA Stream : UBatchContext utilise threading.Barrier + cpu_wait_event/cpu_signal_event + gpu_comm_done_event/gpu_compute_done_event(ubatching.py:25-48) pour synchroniser côté CPU comme côté GPU les deux threads ubatch : yield_and_switch_from_compute_to_comm bascule le stream courant de compute vers comm et attend la fin du compute GPU(ubatching.py:133-139), et inversement.
  • Capture : thread dédié pour init le contexte CUDA : _capture_ubatches lance un thread par ubatch, qui appelle d'abord torch.cuda.current_blas_handle() pour initialiser le contexte CUDA sur ce thread, puis entre dans UBatchContext pour attendre la barrier(gpu_ubatch_wrapper.py:236-251), sans quoi le stream du thread principal et le contexte des threads ubatch ne correspondraient pas à la capture.
  • Partition SM entre communication et calcul : SMControlContextManager(gpu_ubatch_wrapper.py:68-112) réserve via set_comm_sms/set_compute_sms un nombre fixe de SM pour le all2all DeepEP, afin que le all2all massif du MoE ne s'accapare pas tous les SM au détriment du compute.
  • Pool de graphes partagé : UBatchWrapper.graph_pool est transmis à CUDAGraphWrapper.graph_pool(gpu_ubatch_wrapper.py:143-147), tous les graphes partagent le même pool mémoire — on capture les grands d'abord pour que les petits réutilisent(gpu_model_runner.py:6662-6668).

Fichiers clés

  • CUDAGraphMode:53-87 — enum : NONE/PIECEWISE/FULL/FULL_DECODE_ONLY/FULL_AND_PIECEWISE, avec helpers decode_mode/mixed_mode/has_full_cudagraphs.
  • cudagraph_batch_sizes:729-736 — lit depuis compilation_config.cudagraph_capture_sizes les formes enregistrables et les trie.
  • cudagraph_dispatcher:844-845CUDagraphDispatcher(vllm_config) centralise dispatch et capture.
  • _determine_batch_execution_and_padding:3836-3948 — calcule uniform_decode/has_lora/has_encoder_output, appelle cudagraph_dispatcher.dispatch(num_tokens, ...) pour choisir le mode et padder num_tokens jusqu'à BatchDescriptor.num_tokens.
  • maybe_create_ubatch_slices:4197-4203 — quand should_ubatch=True, découpe selon num_tokens_padded // num_ubatches, renvoie (ubatch_slices, ubatch_slices_padded) en deux exemplaires, le padded étant réservé à l'attention.
  • load_model wrapper selection:5350-5379 — selon cudagraph_mode + use_ubatching, enveloppe self.model de BreakableCUDAGraphWrapper / CUDAGraphWrapper / UBatchWrapper.
  • capture_model:6647-6712 — capture les grands shapes en premier, set_cudagraph_capturing_enabled(True), passe par cudagraph_dispatcher.get_capture_descs(), puis lock_workspace().
  • UBatchSlice:13-29 — dataclass request_slice + token_slice, avec is_empty() et la propriété num_tokens.
  • check_ubatch_thresholds:38-46 — si use_ubatching est off, renvoie False directement ; uniform decode utilise dbo_decode_token_threshold, sinon dbo_prefill_token_threshold.
  • maybe_create_ubatch_slices:63-114 — via searchsorted sur cu_num_tokens, trouve le request_slice de chaque ubatch, puis _pad_out_ubatch_slices pad la queue jusqu'à num_tokens_padded.
  • UBatchContext:20-148 — coordinateur d'événements CPU+GPU ; yield_and_switch_from_compute_to_comm/yield_and_switch_from_comm_to_compute sont les cœurs(ubatching.py:133-147).
  • dbo_enabled / dbo_current_ubatch_id:150-157 — via un dictionnaire global _THREAD_ID_TO_CONTEXT, retrouve à quel ubatch appartient le thread courant ; sert aux kernels pour sélectionner le bon buffer.
  • UBatchWrapper.__init__:113-150 — crée comm_stream, ready_barrier(num_ubatches+1), cudagraphs: dict[int, CUDAGraphMetaData] ; selon runtime_mode, décide de construire ou non un CUDAGraphWrapper.
  • _capture_ubatches:212-303 — lance les threads ubatch pour initialiser le contexte CUDA, le thread principal enveloppe dans torch.cuda.graph(...) et join tous les threads ; après capture, self.cudagraphs[num_tokens] = cudagraph_metadata.
  • _run_ubatches:305-341 — chemin de repli sans cudagraph, exécute le modèle en multi-thread pur.
  • UBatchWrapper.__call__:441-537 — récupère ubatch_slices/cudagraph_runtime_mode depuis forward_context ; en FULL et si num_tokens in self.cudagraphs, fait cudagraph_metadata.cudagraph.replay()(gpu_ubatch_wrapper.py:516-521), sinon capture ou _run_ubatches.
  • CUDAGraphWrapper class:145-233concrete_cudagraph_entries: dict[BatchDescriptor, CUDAGraphEntry] ; runtime_mode distingue FULL/PIECEWISE.
  • CUDAGraphWrapper.__call__:233-361 — au premier BatchDescriptor, enregistre via torch.cuda.graph(...) ; ensuite entry.cudagraph.replay(), après get_offloader().sync_prev_onload() pour attendre le prefetch de l'offloader.
  • BreakableCUDAGraphWrapper:246 — implémentation du mode PIECEWISE ; le segment attention est sorti en eager.

Flux de données

Le execute_model du runner atteint _determine_batch_execution_and_padding pour choisir cudagraph_mode et batch_descriptor, puis maybe_create_ubatch_slices découpe les micro-batchs, enfin set_forward_context injecte ces valeurs dans le forward context ; l'appel self.model(...) entre dans le wrapper :

python
# vllm/v1/worker/gpu_model_runner.py L4197-L4203
ubatch_slices, ubatch_slices_padded = maybe_create_ubatch_slices(
    should_ubatch,
    num_scheduled_tokens_np,
    num_tokens_padded,
    num_reqs_padded,
    self.parallel_config.num_ubatches,
)

Une fois entré dans UBatchWrapper.__call__, on regarde si forward_context.ubatch_slices est None. Si oui, chemin mono-graphe (ou eager simple) ; sinon, chemin multi-thread ubatch :

python
# vllm/v1/worker/gpu_ubatch_wrapper.py L493-L521
if (
    num_tokens not in self.cudagraphs
    and cudagraph_runtime_mode is CUDAGraphMode.FULL
):
    ubatch_metadata = self._make_ubatch_metadata(
        ubatch_slices=ubatch_slices,
        attn_metadata=attn_metadata,
        slot_mapping=slot_mapping,
        input_ids=input_ids,
        positions=positions,
        intermediate_tensors=intermediate_tensors,
        inputs_embeds=inputs_embeds,
        compute_stream=compute_stream,
        dp_metadata=ubatch_dp_metadata,
        batch_descriptor=batch_descriptor,
        cudagraph_runtime_mode=CUDAGraphMode.NONE,
    )
    with self.sm_control:
        return self._capture_ubatches(ubatch_metadata, self.runnable)
elif (
    num_tokens in self.cudagraphs
    and cudagraph_runtime_mode is CUDAGraphMode.FULL
):
    cudagraph_metadata = self.cudagraphs[num_tokens]
    get_offloader().sync_prev_onload()
    cudagraph_metadata.cudagraph.replay()
    return cudagraph_metadata.outputs

_capture_ubatches lance N threads, chacun entre dans son propre UBatchContext ; côté CPU, cpu_wait_event/cpu_signal_event font alterner l'un court et l'autre attend, côté GPU gpu_comm_done_event/gpu_compute_done_event record + wait sur comm_stream et compute_stream(ubatching.py:82-92) ; au final, le thread principal enregistre le tout dans le contexte torch.cuda.graph(...) et le stocke dans self.cudagraphs[num_tokens]. Ensuite, même forme → replay(), sans overhead Python.

Limites et échecs

  • Repli sur batch vide : maybe_create_ubatch_slices renvoie directement None, None quand should_ubatch=False(ubatch_utils.py:71-72) ; UBatchWrapper.__call__, voyant ubatch_slices is None, prend le chemin mono-graphe(gpu_ubatch_wrapper.py:447-464).
  • Conflit de formes en FULL : les shapes qui activent ubatching ne sont enregistrées que pour le graphe ubatch ; pour éviter qu'une même forme sans ubatch ne soit enregistrée deux fois par CUDAGraphWrapper, __call__ force explicitement if batch_descriptor.num_tokens in self.cudagraphs: cudagraph_runtime_mode = NONE(gpu_ubatch_wrapper.py:455-458).
  • DP incohérent → assertion : dans le chemin ubatch, UBatchWrapper.__call__ fait assert dp_metadata is not None(gpu_ubatch_wrapper.py:477-478) ; ubatch est conçu pour DP>1, un DP unique ne peut pas y entrer.
  • Dernier ubatch vide : is_last_ubatch_empty vérifie padded_num_tokens // num_ubatches * (num_ubatches-1) >= orig_num_tokens(ubatch_utils.py:32-35) ; dans ce cas, le dernier ubatch n'a pas de token réel et doit être skippé.
  • Sync offloader pendant la capture : _capture_ubatches appelle get_offloader().sync_prev_onload() avant la capture(gpu_ubatch_wrapper.py:283-285) et get_offloader().join_after_forward() après(gpu_ubatch_wrapper.py:298-301), sinon le stream de prefetch de l'offloader est déconnecté et remonte une erreur unjoined stream.
  • cudagraph n'autorise pas la capture accidentelle à l'exécution : set_cudagraph_capturing_enabled(False) à la fin de capture_model(gpu_model_runner.py:6689-6694), et validate_cudagraph_capturing_enabled est vérifié avant la capture dans CUDAGraphWrapper(cuda_graph.py:276-277) ; un chemin de capture inattendu à l'exécution lève une erreur.
  • Synchronisation de forme entre DP ranks : coordinate_batch_across_dp décide should_ubatch et num_tokens_across_dp entre ranks(gpu_model_runner.py:3907-3918) ; tous les ranks découpent les mêmes formes d'ubatch pour que le all2all ne se désaligne pas au replay.
  • Stabilité des adresses d'input : CUDAGraphWrapper en mode debug fait assert new_input_addresses == entry.input_addresses(cuda_graph.py:346-355) ; les persistent_buffers sont préparés par le runner, sinon le replay lirait d'anciennes adresses.

Résumé

cudagraph est le levier principal qui fait décoller decode et prefill en v1 : forme fixe, un enregistrement pour N replays ; CUDAGraphWrapper utilise BatchDescriptor comme clé, CUDAGraphDispatcher décide du mode par batch. ubatch découpe un batch en N petits de forme égale quand DP>1, pour chevaucher communication et calcul entre DP ranks, via les événements CPU + GPU stream event de UBatchContext qui alternent entre deux threads. L'ensemble forme le DBO, presque indispensable dans les déploiements massifs de MoE. Pour la manière dont le forward atteint cette couche, voir /worker/gpu-model-runner ; pour la manière dont le worker appelle le runner, voir /worker/worker ; pour la couche exécuteur, voir /executor/executor.

Voir la documentation officielle : Documentation vLLM · README