Skip to content

ubatch y cudagraph: cómo los grafos de forma fija aceleran decode/prefill

源码版本v0.25.1

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__ usa BatchDescriptor como clave (cuda_graph.py:257-263); cada forma se graba una sola vez y la siguiente vez se hace replay directamente. El CUDAGraphDispatcher centraliza en el runner "qué formas se graban con qué modo" (gpu_model_runner.py:844-845); _determine_batch_execution_and_padding hace pad de num_tokens a 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_slices parte por num_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: UBatchContext usa threading.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_comm cambia 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_ubatches lanza un hilo por ubatch, primero llama a torch.cuda.current_blas_handle() para que el CUDA context se inicialice en ese hilo, y luego entra al UBatchContext a 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ía set_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_pool se reenvía a CUDAGraphWrapper.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 como decode_mode/mixed_mode/has_full_cudagraphs.
  • cudagraph_batch_sizes:729-736 — lee y ordena las formas grabables desde compilation_config.cudagraph_capture_sizes.
  • cudagraph_dispatcher:844-845CUDagraphDispatcher(vllm_config) centraliza el dispatch y la captura.
  • _determine_batch_execution_and_padding:3836-3948 — calcula uniform_decode/has_lora/has_encoder_output, llama a cudagraph_dispatcher.dispatch(num_tokens, ...) para elegir el modo y hace pad de num_tokens al BatchDescriptor.num_tokens.
  • maybe_create_ubatch_slices:4197-4203 — con should_ubatch=True corta por num_tokens_padded // num_ubatches y devuelve dos copias, (ubatch_slices, ubatch_slices_padded); la versión padded es para attention por separado.
  • load_model wrapper selection:5350-5379 — según cudagraph_mode + use_ubatching envuelve self.model con BreakableCUDAGraphWrapper / CUDAGraphWrapper / UBatchWrapper.
  • capture_model:6647-6712 — captura primero las formas grandes, set_cudagraph_capturing_enabled(True), recorre cudagraph_dispatcher.get_capture_descs() y al final lock_workspace().
  • UBatchSlice:13-29 — dataclass con request_slice + token_slice; propiedades is_empty() y num_tokens.
  • check_ubatch_thresholds:38-46 — si use_ubatching está apagado devuelve False directamente; para uniform decode usa dbo_decode_token_threshold, en caso contrario dbo_prefill_token_threshold.
  • maybe_create_ubatch_slices:63-114 — usa searchsorted sobre cu_num_tokens para encontrar el request_slice de cada ubatch y por último _pad_out_ubatch_slices rellena la cola hasta num_tokens_padded.
  • UBatchContext:20-148 — coordinador de eventos CPU+GPU; yield_and_switch_from_compute_to_comm/yield_and_switch_from_comm_to_compute son el núcleo (ubatching.py:133-147).
  • dbo_enabled / dbo_current_ubatch_id:150-157 — consulta el diccionario global _THREAD_ID_TO_CONTEXT para 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 — crea comm_stream, ready_barrier(num_ubatches+1), cudagraphs: dict[int, CUDAGraphMetaData] y según runtime_mode decide si crear el CUDAGraphWrapper.
  • _capture_ubatches:212-303 — lanza hilos ubatch para inicializar el CUDA context; el hilo principal envuelve con torch.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-537forward_context obtiene ubatch_slices/cudagraph_runtime_mode; en modo FULL y con num_tokens in self.cudagraphs hace cudagraph_metadata.cudagraph.replay() (gpu_ubatch_wrapper.py:516-521); en caso contrario captura o llama a _run_ubatches.
  • CUDAGraphWrapper class:145-233concrete_cudagraph_entries: dict[BatchDescriptor, CUDAGraphEntry]; runtime_mode distingue FULL/PIECEWISE.
  • CUDAGraphWrapper.__call__:233-361 — la primera vez que entra un BatchDescriptor se graba el grafo con torch.cuda.graph(...); después entry.cudagraph.replay(); antes del replay get_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:

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,
)

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:

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 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_slices con should_ubatch=False devuelve directamente None, None (ubatch_utils.py:71-72); al ver ubatch_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ícitamente if 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__ hace assert 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_empty comprueba padded_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_ubatches antes de capturar llama a get_offloader().sync_prev_onload() (gpu_ubatch_wrapper.py:283-285) y después a get_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 de capture_model (gpu_model_runner.py:6689-6694); validate_cudagraph_capturing_enabled hace el check antes de la captura en CUDAGraphWrapper (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_dp decide entre ranks el should_ubatch y el num_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, CUDAGraphWrapper hace assert new_input_addresses == entry.input_addresses (cuda_graph.py:346-355); los persistent_buffers los 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.

Véase la documentación oficial: vLLM 文档 · README