Skip to content

parallel_state y GroupCoordinator: grupos de comunicación TP/PP/DP

源码版本v0.25.1

Responsabilidades

vLLM ejecuta inferencia distribuida apoyándose en torch.distributed de PyTorch mediante ProcessGroup, pero el ProcessGroup crudo es verboso y difícil de unificar: la comunicación por CPU y por device exigen dos groups distintos, el all-reduce personalizado debe pasar por una custom op, y la difusión de objetos frente a la de tensores utiliza dos APIs diferentes. GroupCoordinator(parallel_state.py:358) es una capa de envoltura que vLLM añade sobre ProcessGroup: empaqueta en un solo objeto el cpu group, el device group, el device_communicator y el mq_broadcaster de un conjunto de procesos, ofreciendo una interfaz unificada con all_reduce / all_gather / broadcast / send_object / recv_tensor_dict / barrier.

El módulo parallel_state mantiene varios singletons globales: _TP(paralelismo de tensores), _PP(paralelismo de pipeline), _DP(paralelismo de datos), _EP(paralelismo de expertos), _DCP(decode context parallel), _PCP(prefill context parallel), _EPLB(equilibrado de carga de expertos). Cada group se obtiene mediante get_tp_group / get_pp_group / get_dp_group etc.(parallel_state.py:1368-1424); las capas del modelo como ColumnParallelLinear / RowParallelLinear dependen de estos getters para obtener su grupo de comunicación. La inicialización de los grupos la realiza initialize_model_parallel(parallel_state.py:1713) una vez al arranque, repartiendo el rank global en segmentos según tensor_model_parallel_size * pipeline_model_parallel_size * ... y llamando new_group.

GroupCoordinator además hace una cosa más: registra all_reduce / all_gather como custom ops torch.ops.vllm.*, de modo que Dynamo al compilar pueda fusionar la op de comunicación dentro del graph como una op cualquiera(parallel_state.py:641-663), en lugar de meter un objeto Python como self en el graph de compilación. Cuando el world size es 1, todas las ops de comunicación hacen directamente return input y saltan NCCL.

Motivación de diseño

  • Un objeto, dos PGs: cpu_group(backend gloo) y device_group(nccl/gloo/backend de plataforma) se separan(parallel_state.py:435-452); la difusión de objetos y la coordinación entre coordinadores ocurre en CPU, la comunicación de tensores en device.
  • Tres conceptos de rank bien diferenciados: rank(global), local_rank(dentro del nodo), rank_in_group(dentro del grupo)(parallel_state.py:369-380) — en escenarios multinode estos tres difieren, y la capa del modelo debe usar rank_in_group para partir correctamente las particiones.
  • Layout de ranks con orden fijo: ExternalDP × DP × PP × PCP × TP(parallel_state.py:1779-1794); tras torch.arange(world_size).reshape(...) se hace unbind por una dimensión para obtener la lista de ranks de cada grupo.
  • Bypass en mono-GPU: métodos como all_reduce hacen return directo cuando world_size == 1(parallel_state.py:657-658), evitando arrancar NCCL aunque haya una sola GPU.
  • Custom op para que Dynamo no proteste: con use_custom_op_call=True se llama torch.ops.vllm.all_reduce(input_, group_name=self.unique_name)(parallel_state.py:660-661); el grupo se pasa como string y se resuelve por tabla, para que Dynamo no se tropiece con un objeto Python.
  • DeviceCommunicator enchufable: el campo device_communicator permite implementaciones CUDA / XPU / CPU / all-reduce personalizado(parallel_state.py:463-478); si is_cuda_aliked se vincula a cuda:N, XPU se vincula a xpu:N.
  • TP emite con cola de mensajes: init_model_parallel_group abre use_message_queue_broadcaster=True en el grupo TP(parallel_state.py:1805-1811), para que los workers compartan diccionarios de tensores sin pasar por NCCL.
  • StatelessGroupCoordinator como red: en escenarios EP elásticos sin torch.distributed inicializado, _init_stateless_group(parallel_state.py:1316-1338) levanta un grupo equivalente con TCP store + StatelessProcessGroup.

Archivos clave

Flujo de datos

Al arrancar, Worker llama initialize_model_parallel(tp_size, pp_size, ...) dentro de su propio proceso; esta función hace reshape del rank global con el orden ExternalDP × DP × PP × PCP × TP y luego unbind por cada dimensión para obtener la lista de ranks de cada grupo. La construcción del grupo TP es:

python
# vllm/distributed/parallel_state.py L1788-L1811
all_ranks = torch.arange(world_size).reshape(
    -1,
    data_parallel_size,
    pipeline_model_parallel_size,
    prefill_context_model_parallel_size,
    tensor_model_parallel_size,
)  # noqa

# Build the tensor model-parallel groups.
global _TP
assert _TP is None, "tensor model parallel group is already initialized"
group_ranks = all_ranks.view(-1, tensor_model_parallel_size).unbind(0)
group_ranks = [x.tolist() for x in group_ranks]
if enable_elastic_ep:
    group_ranks = local_all_ranks.view(-1, tensor_model_parallel_size).unbind(0)
    group_ranks = [x.tolist() for x in group_ranks]
# message queue broadcaster is only used in tensor model parallel group
_TP = init_model_parallel_group(
    group_ranks,
    get_world_group().local_rank,
    backend,
    use_message_queue_broadcaster=True,
    group_name="tp",
)

Cuando la capa del modelo invoca comunicación no toca ProcessGroup directamente, sino get_tp_group().all_reduce(tensor); dentro, all_reduce decide según use_custom_op_call si va por la custom op o si llama directamente a device_communicator.all_reduce:

python
# vllm/distributed/parallel_state.py L641-L668
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
    """
    User-facing all-reduce function before we actually call the
    all-reduce operation.
    """
    # Bypass the function if we are using only 1 GPU.
    if self.world_size == 1:
        return input_

    if self.use_custom_op_call:
        return torch.ops.vllm.all_reduce(input_, group_name=self.unique_name)
    else:
        return self._all_reduce_out_place(input_)

def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
    if self.device_communicator is None:
        raise ValueError("No device communicator found")
    return self.device_communicator.all_reduce(input_)

El forward de ColumnParallelLinear, tras computar la parte local, llama all_reduce para fusionar las particiones; RowParallelLinear hace lo opuesto, primero all-gather, luego la parte local, y al final reduce_scatter. Todas estas ops quedan unificadas en la capa GroupCoordinator.

Límites y fallos

  • rank_in_group debe usarse con el rank interno del grupo: la capa del modelo debe partir con rank_in_group y no con el rank global, de lo contrario en escenarios multinode cortará particiones equivocadas(parallel_state.py:380).
  • Inicializar TP dos veces rompe el assert: initialize_model_parallel verifica assert _TP is None para cada grupo(parallel_state.py:1798); llamarlo de nuevo lanza AssertionError, antes hay que pasar por destroy_model_parallel.
  • DCP no puede superar a TP: dcp_size <= tp_size, porque DCP reutiliza las GPU del grupo TP(parallel_state.py:1816-1820).
  • Monoplaca hace bypass: all_reduce con world_size == 1 devuelve return input(parallel_state.py:657-658), sin necesidad de proceso NCCL; pero si device_communicator is None y no aplica bypass, lanza ValueError.
  • DeviceCommunicator ausente lanza error: _all_reduce_out_place verifica device_communicator is None y lanza(parallel_state.py:665-668), para evitar que se invoque sin backend inicializado.
  • _replace_active_groups debe llamarse colectivamente: parallel_state.py:1341-1362 exige que todos los ranks lo invoquen a la vez, de lo contrario deadlock; en escalado elástico de EP lo orquesta el supervisor.

Resumen

parallel_state es la base de la inferencia multigpu de vLLM: GroupCoordinator envuelve PyTorch ProcessGroup, juntando en un mismo objeto el dual group CPU/device, la custom op, el device_communicator y el broadcaster por cola de mensajes; la capa del modelo obtiene su grupo con get_tp_group() / get_pp_group() y luego invoca all_reduce y similares. Cómo usan estas comunicaciones las capas lineales TP, en /model-loading/tp-layers;cómo arranca el Worker estos grupos, en /worker/worker;el núcleo del motor que engancha a los workers, en /engine/engine-core.

Véase la documentación oficial: documentación de vLLM · README.