ubatch y cudagraph: cómo los grafos de forma fija aceleran decode/prefill
Responsabilidades
cudagraph es el mecanismo de CUDA de "grabar una vez, reproducir infinitas veces": se registra una secuencia de kernels como un grafo y en cada replay se ahorra el overhead de lanzamiento. En decode, donde cada token solo calcula uno o dos tokens nuevos y los kernels son pequeños y numerosos, es justo el terreno de cudagraph. v1 usa CUDAGraphWrapper (cuda_graph.py:145-233) para envolver el forward del modelo; en __call__, la primera vez que aparece un BatchDescriptor se graba el grafo (cuda_graph.py:265-344) y para las siguientes llamadas con la misma forma se hace directamente entry.cudagraph.replay() (cuda_graph.py:357-361).
ubatch (micro-batch) es la forma en que v1, cuando DP>1, corta un step de batch grande en N pequeños de forma fija; cada micro-batch se graba como su propio grafo, de modo que distintos ranks DP pueden correr a la vez, hacer la comunicación TP/EP all2all en paralelo y solapar comunicación con cómputo. Esta lógica está implementada en UBatchWrapper (gpu_ubatch_wrapper.py:113-211) y UBatchContext (ubatching.py:20-148). Juntos forman DBO (Distributed Batch Overlap): partición en micro-batches + captura de cudagraph + coordinación con eventos multi-hilo.
El modo de ejecución de cudagraph lo decide CUDAGraphMode (compilation.py:53-87): NONE (no grabar), PIECEWISE (captura por segmentos, attention no entra al grafo), FULL (captura la totalidad), FULL_DECODE_ONLY (solo en decode se graba full; prefill va por piecewise), etc. Al final de GPUModelRunner.load_model, según la configuración, se envuelve self.model con el wrapper correspondiente (gpu_model_runner.py:5350-5379).
Motivación de diseño
- cudagraph dirigido por la forma:
CUDAGraphWrapper.__call__usaBatchDescriptorcomo clave (cuda_graph.py:257-263); cada forma se graba una sola vez y la siguiente vez se hace replay directamente. ElCUDAGraphDispatchercentraliza en el runner "qué formas se graban con qué modo" (gpu_model_runner.py:844-845);_determine_batch_execution_and_paddinghace pad denum_tokensa la forma grabable más cercana (gpu_model_runner.py:3879-3894). - PIECEWISE salta attention: en un grafo full, los kernels de attention tienen formas demasiado flexibles (query/key de longitud variable); si se graban, el siguiente replay no serviría. v1 ofrece el modo piecewise, que corta el modelo en varios segmentos, deja attention en eager y mete el resto en el grafo, implementado por
BreakableCUDAGraphWrapper(breakable_cudagraph.py:246). - Partición uniforme del ubatch:
maybe_create_ubatch_slicesparte pornum_tokens_padded // num_ubatches(ubatch_utils.py:63-114), de modo que cada ubatch tenga la misma forma; así el mismo cudagraph se puede replay varias veces y alinearse entre ranks DP distintos. - Coordinación multi-hilo con eventos de CUDA Stream:
UBatchContextusathreading.Barrier+cpu_wait_event/cpu_signal_event+gpu_comm_done_event/gpu_compute_done_event(ubatching.py:25-48) para sincronizar en lado CPU y GPU los dos hilos ubatch:yield_and_switch_from_compute_to_commcambia el stream actual de compute a comm y espera a que el cómputo GPU termine (ubatching.py:133-139), y al revés. - En capture, un hilo aparte inicializa el CUDA context:
_capture_ubatcheslanza un hilo por ubatch, primero llama atorch.cuda.current_blas_handle()para que el CUDA context se inicialice en ese hilo, y luego entra alUBatchContexta esperar la barrier (gpu_ubatch_wrapper.py:236-251); de lo contrario, el stream del hilo principal y el context del hilo ubatch no casarían en la captura. - Partición de SMs entre comunicación y cómputo:
SMControlContextManager(gpu_ubatch_wrapper.py:68-112) reserva, víaset_comm_sms/set_compute_sms, una cantidad fija de SMs para la comunicación all2all de DeepEP, evitando que el all2all masivo de MoE ocupe toda la GPU y que el cómputo se quede sin SM. - Pool de grafos compartido:
UBatchWrapper.graph_poolse reenvía aCUDAGraphWrapper.graph_pool(gpu_ubatch_wrapper.py:143-147); todos los grafos de todas las formas comparten el mismo pool de memoria, los grandes se capturan primero y los pequeños lo reutilizan (gpu_model_runner.py:6662-6668).
Archivos clave
CUDAGraphMode:53-87— enumerado:NONE/PIECEWISE/FULL/FULL_DECODE_ONLY/FULL_AND_PIECEWISE; helpers comodecode_mode/mixed_mode/has_full_cudagraphs.cudagraph_batch_sizes:729-736— lee y ordena las formas grabables desdecompilation_config.cudagraph_capture_sizes.cudagraph_dispatcher:844-845—CUDagraphDispatcher(vllm_config)centraliza el dispatch y la captura._determine_batch_execution_and_padding:3836-3948— calculauniform_decode/has_lora/has_encoder_output, llama acudagraph_dispatcher.dispatch(num_tokens, ...)para elegir el modo y hace pad denum_tokensalBatchDescriptor.num_tokens.maybe_create_ubatch_slices:4197-4203— conshould_ubatch=Truecorta pornum_tokens_padded // num_ubatchesy devuelve dos copias,(ubatch_slices, ubatch_slices_padded); la versión padded es para attention por separado.load_model wrapper selection:5350-5379— segúncudagraph_mode+use_ubatchingenvuelveself.modelconBreakableCUDAGraphWrapper/CUDAGraphWrapper/UBatchWrapper.capture_model:6647-6712— captura primero las formas grandes,set_cudagraph_capturing_enabled(True), recorrecudagraph_dispatcher.get_capture_descs()y al finallock_workspace().UBatchSlice:13-29— dataclass conrequest_slice+token_slice; propiedadesis_empty()ynum_tokens.check_ubatch_thresholds:38-46— siuse_ubatchingestá apagado devuelve False directamente; para uniform decode usadbo_decode_token_threshold, en caso contrariodbo_prefill_token_threshold.maybe_create_ubatch_slices:63-114— usasearchsortedsobrecu_num_tokenspara encontrar elrequest_slicede cada ubatch y por último_pad_out_ubatch_slicesrellena la cola hastanum_tokens_padded.UBatchContext:20-148— coordinador de eventos CPU+GPU;yield_and_switch_from_compute_to_comm/yield_and_switch_from_comm_to_computeson el núcleo (ubatching.py:133-147).dbo_enabled / dbo_current_ubatch_id:150-157— consulta el diccionario global_THREAD_ID_TO_CONTEXTpara saber en qué ubatch está el hilo actual; el código lo usa para que el kernel elija el buffer del ubatch actual.UBatchWrapper.__init__:113-150— creacomm_stream,ready_barrier(num_ubatches+1),cudagraphs: dict[int, CUDAGraphMetaData]y segúnruntime_modedecide si crear elCUDAGraphWrapper._capture_ubatches:212-303— lanza hilos ubatch para inicializar el CUDA context; el hilo principal envuelve contorch.cuda.graph(...)el join de todos los hilos y, tras la captura,self.cudagraphs[num_tokens] = cudagraph_metadata._run_ubatches:305-341— ruta de fallback sin cudagraph; corre el modelo en puro multi-hilo.UBatchWrapper.__call__:441-537—forward_contextobtieneubatch_slices/cudagraph_runtime_mode; en modo FULL y connum_tokens in self.cudagraphshacecudagraph_metadata.cudagraph.replay()(gpu_ubatch_wrapper.py:516-521); en caso contrario captura o llama a_run_ubatches.CUDAGraphWrapper class:145-233—concrete_cudagraph_entries: dict[BatchDescriptor, CUDAGraphEntry];runtime_modedistingue FULL/PIECEWISE.CUDAGraphWrapper.__call__:233-361— la primera vez que entra unBatchDescriptorse graba el grafo contorch.cuda.graph(...); despuésentry.cudagraph.replay(); antes del replayget_offloader().sync_prev_onload()espera a que el offloader complete el prefetch.BreakableCUDAGraphWrapper:246— implementación del modo PIECEWISE; el segmento de attention se corta y se ejecuta en eager.
Flujo de datos
Cuando el execute_model del runner llega a _determine_batch_execution_and_padding, elige cudagraph_mode y batch_descriptor, luego maybe_create_ubatch_slices corta los micro-batches y por último set_forward_context vuelca todo en el forward context; basta con llamar self.model(...) para entrar al 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,
)Al entrar a UBatchWrapper.__call__, comprueba si forward_context.ubatch_slices es None. Si lo es, va por la ruta de un solo grafo (o plain eager); si no, va por la ruta multi-hilo de 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 lanza N hilos, cada uno entra a su propio UBatchContext; en el lado CPU los cpu_wait_event/cpu_signal_event hacen que se alternen: uno corre y otro espera; en el lado GPU, los gpu_comm_done_event/gpu_compute_done_event hacen record + wait sobre comm_stream y compute_stream (ubatching.py:82-92). Finalmente el hilo principal, dentro del contexto torch.cuda.graph(...), graba todo el tramo como un único grafo almacenado en self.cudagraphs[num_tokens]. A partir de ahí, una misma forma entra directamente por replay(), sin overhead de Python.
Límites y fallos
- Fallback en batch vacío:
maybe_create_ubatch_slicesconshould_ubatch=Falsedevuelve directamenteNone, None(ubatch_utils.py:71-72); al verubatch_slices is None,UBatchWrapper.__call__va por la ruta de un solo grafo (gpu_ubatch_wrapper.py:447-464). - Conflicto de formas en modo full: una forma con ubatching activado solo graba el grafo ubatch; para evitar que la misma forma sin ubatch se grabe dos veces por
CUDAGraphWrapper,__call__hace explícitamenteif batch_descriptor.num_tokens in self.cudagraphs: cudagraph_runtime_mode = NONE(gpu_ubatch_wrapper.py:455-458). - DP inconsistente, assert directo: en la ruta ubatch,
UBatchWrapper.__call__haceassert dp_metadata is not None(gpu_ubatch_wrapper.py:477-478); ubatch está diseñado para DP>1, con un único DP no se llega aquí. - El último ubatch vacío:
is_last_ubatch_emptycompruebapadded_num_tokens // num_ubatches * (num_ubatches-1) >= orig_num_tokens(ubatch_utils.py:32-35); en ese caso el último ubatch no tiene tokens reales y el código lo salta. - Sincronización del offloader durante la captura:
_capture_ubatchesantes de capturar llama aget_offloader().sync_prev_onload()(gpu_ubatch_wrapper.py:283-285) y después aget_offloader().join_after_forward()(gpu_ubatch_wrapper.py:298-301); si no, el stream de prefetch del offloader no se engancha y aparece un error de stream unjoined. - cudagraph no permite capturas accidentales en runtime:
set_cudagraph_capturing_enabled(False)se ejecuta al final decapture_model(gpu_model_runner.py:6689-6694);validate_cudagraph_capturing_enabledhace el check antes de la captura enCUDAGraphWrapper(cuda_graph.py:276-277) y entrar por la ruta de captura de forma imprevista en runtime lanza un error. - Sincronización de formas entre ranks DP:
coordinate_batch_across_dpdecide entre ranks elshould_ubatchy elnum_tokens_across_dp(gpu_model_runner.py:3907-3918) para que todos los ranks corten ubatches con la misma forma; en otro caso el replay del all2all quedaría desalineado. - Las direcciones de input deben ser estables: en modo debug,
CUDAGraphWrapperhaceassert new_input_addresses == entry.input_addresses(cuda_graph.py:346-355); lospersistent_bufferslos prepara el runner; si no, el replay leería direcciones viejas.
Resumen
cudagraph es la herramienta central con la que v1 hace volar decode y prefill: graba una vez para una forma fija y luego replay repetidamente. CUDAGraphWrapper usa BatchDescriptor como clave y CUDAGraphDispatcher decide qué modo le toca a cada batch. ubatch, cuando DP>1, corta un batch en N micro-batches de forma idéntica para que la comunicación y el cómputo entre ranks DP se solapen; se apoya en los eventos CPU + eventos de stream GPU de UBatchContext para alternar entre dos hilos. La unión de ambos es DBO, casi obligatorio en despliegues grandes de MoE. Para cómo el forward llega a esta capa, véase /worker/gpu-model-runner; para cómo el worker invoca al runner, /worker/worker; la capa del executor, /executor/executor.