ubatch et cudagraph : comment les graphes à forme fixe accélèrent decode/prefill
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__utiliseBatchDescriptorcomme clé(cuda_graph.py:257-263), chaque forme n'est enregistrée qu'une fois, puis on rejoue ;CUDAGraphDispatchercentralise dans le runner « quel mode pour quelle forme »(gpu_model_runner.py:844-845) ;_determine_batch_execution_and_paddingpadnum_tokensjusqu'à 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_slicesdécoupe au pointnum_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 :
UBatchContextutilisethreading.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_commbascule 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_ubatcheslance un thread par ubatch, qui appelle d'abordtorch.cuda.current_blas_handle()pour initialiser le contexte CUDA sur ce thread, puis entre dansUBatchContextpour 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 viaset_comm_sms/set_compute_smsun 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_poolest 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 helpersdecode_mode/mixed_mode/has_full_cudagraphs.cudagraph_batch_sizes:729-736— lit depuiscompilation_config.cudagraph_capture_sizesles formes enregistrables et les trie.cudagraph_dispatcher:844-845—CUDagraphDispatcher(vllm_config)centralise dispatch et capture._determine_batch_execution_and_padding:3836-3948— calculeuniform_decode/has_lora/has_encoder_output, appellecudagraph_dispatcher.dispatch(num_tokens, ...)pour choisir le mode et paddernum_tokensjusqu'àBatchDescriptor.num_tokens.maybe_create_ubatch_slices:4197-4203— quandshould_ubatch=True, découpe selonnum_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— seloncudagraph_mode+use_ubatching, enveloppeself.modeldeBreakableCUDAGraphWrapper/CUDAGraphWrapper/UBatchWrapper.capture_model:6647-6712— capture les grands shapes en premier,set_cudagraph_capturing_enabled(True), passe parcudagraph_dispatcher.get_capture_descs(), puislock_workspace().UBatchSlice:13-29— dataclassrequest_slice+token_slice, avecis_empty()et la propriéténum_tokens.check_ubatch_thresholds:38-46— siuse_ubatchingest off, renvoie False directement ; uniform decode utilisedbo_decode_token_threshold, sinondbo_prefill_token_threshold.maybe_create_ubatch_slices:63-114— viasearchsortedsurcu_num_tokens, trouve lerequest_slicede chaque ubatch, puis_pad_out_ubatch_slicespad 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_computesont 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éecomm_stream,ready_barrier(num_ubatches+1),cudagraphs: dict[int, CUDAGraphMetaData]; selonruntime_mode, décide de construire ou non unCUDAGraphWrapper._capture_ubatches:212-303— lance les threads ubatch pour initialiser le contexte CUDA, le thread principal enveloppe danstorch.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èreubatch_slices/cudagraph_runtime_modedepuisforward_context; en FULL et sinum_tokens in self.cudagraphs, faitcudagraph_metadata.cudagraph.replay()(gpu_ubatch_wrapper.py:516-521), sinon capture ou_run_ubatches.CUDAGraphWrapper class:145-233—concrete_cudagraph_entries: dict[BatchDescriptor, CUDAGraphEntry];runtime_modedistingue FULL/PIECEWISE.CUDAGraphWrapper.__call__:233-361— au premierBatchDescriptor, enregistre viatorch.cuda.graph(...); ensuiteentry.cudagraph.replay(), aprèsget_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 :
# 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 :
# 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_slicesrenvoie directementNone, Nonequandshould_ubatch=False(ubatch_utils.py:71-72) ;UBatchWrapper.__call__, voyantubatch_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 explicitementif 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__faitassert 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_emptyvérifiepadded_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_ubatchesappelleget_offloader().sync_prev_onload()avant la capture(gpu_ubatch_wrapper.py:283-285) etget_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 decapture_model(gpu_model_runner.py:6689-6694), etvalidate_cudagraph_capturing_enabledest vérifié avant la capture dansCUDAGraphWrapper(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_dpdécideshould_ubatchetnum_tokens_across_dpentre 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 :
CUDAGraphWrapperen mode debug faitassert new_input_addresses == entry.input_addresses(cuda_graph.py:346-355) ; lespersistent_bufferssont 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