Skip to content

parallel_state et GroupCoordinator : groupes de communication TP/PP/DP

源码版本v0.25.1

Responsabilités

L'inférence distribuée de vLLM s'appuie sur les ProcessGroup de torch.distributed de PyTorch, mais les ProcessGroup bruts sont à la fois verbeux et difficiles à unifier : il faut un groupe séparé pour les communications CPU et device, les all-reduce personnalisés passent par des custom ops, et la diffusion d'objets et de tenseurs utilise deux API distinctes. GroupCoordinator(parallel_state.py:358) est la couche d'abstraction que vLLM ajoute par-dessus ProcessGroup : elle regroupe dans un seul objet le cpu_group, le device_group, le device_communicator et le mq_broadcaster d'un ensemble de processus, et expose une interface unifiée all_reduce / all_gather / broadcast / send_object / recv_tensor_dict / barrier.

Le module parallel_state maintient plusieurs singletons globaux : _TP (parallélisme de tenseurs), _PP (parallélisme de pipeline), _DP (parallélisme de données), _EP (parallélisme d'experts), _DCP (decode context parallel), _PCP (prefill context parallel) et _EPLB (équilibrage de charge des experts). Chaque groupe est lu via get_tp_group / get_pp_group / get_dp_group, etc.(parallel_state.py:1368-1424) ; les couches modèle ColumnParallelLinear / RowParallelLinear s'appuient sur ces getters pour récupérer le groupe de communication. L'initialisation des groupes est effectuée une fois au démarrage par initialize_model_parallel(parallel_state.py:1713), qui découpe les rangs globaux selon tensor_model_parallel_size * pipeline_model_parallel_size * ... et appelle new_group.

GroupCoordinator fait aussi une chose : il enregistre all_reduce / all_gather comme des custom ops torch.ops.vllm.*, de sorte que Dynamo puisse les fusionner dans le graphe de compilation comme des ops ordinaires(parallel_state.py:641-663) au lieu d'y injecter l'objet Python self. Quand le world size vaut 1, toutes les ops de communication retournent directement l'entrée et court-circuitent NCCL.

Motivation de conception

  • Un groupe, deux PG : cpu_group (backend gloo) et device_group (backend nccl/gloo/plateforme) sont séparés(parallel_state.py:435-452) : diffusion d'objets et coordination côté CPU, communication de tenseurs côté device.
  • Trois notions de rank bien distinctes : rank (global), local_rank (au sein du nœud), rank_in_group (au sein du groupe)(parallel_state.py:369-380) — en multi-nœud ces trois valeurs diffèrent, et les couches modèle doivent utiliser rank_in_group pour découper correctement.
  • Layout de rangs à ordre fixe : ExternalDP × DP × PP × PCP × TP(parallel_state.py:1779-1794) ; torch.arange(world_size).reshape(...) puis unbind le long d'une dimension donne la liste des rangs du groupe correspondant.
  • Court-circuit mono-GPU : les méthodes comme all_reduce retournent directement quand world_size == 1(parallel_state.py:657-658) pour éviter de démarrer NCCL sur un GPU unique.
  • Custom ops compatibles Dynamo : avec use_custom_op_call=True, on passe par torch.ops.vllm.all_reduce(input_, group_name=self.unique_name)(parallel_state.py:660-661) ; le groupe est passé comme une chaîne et retrouvé dans une table, pour éviter que Dynamo ne bute sur l'objet Python.
  • DeviceCommunicateur pluguable : le champ device_communicator permet à CUDA / XPU / CPU / implémentations personnalisées d'all-reduce d'avoir chacune leur implémentation(parallel_state.py:463-478) ; is_cuda_aliked est lié à cuda:N, XPU à xpu:N.
  • TP diffuse par file de messages : init_model_parallel_group active use_message_queue_broadcaster=True sur le groupe TP(parallel_state.py:1805-1811) pour que les workers partagent des dictionnaires de tenseurs sans passer par NCCL.
  • StatelessGroupCoordinator de secours : pour les scénarios EP élastiques où torch.distributed n'est pas initialisé, _init_stateless_group(parallel_state.py:1316-1338) construit un groupe équivalent à partir d'un TCP store et d'un StatelessProcessGroup.

Fichiers clés

Flux de données

Au démarrage, le Worker appelle initialize_model_parallel(tp_size, pp_size, ...) dans son propre processus ; cette fonction reshape les rangs globaux selon l'ordre ExternalDP × DP × PP × PCP × TP, puis unbind le long de chaque dimension pour obtenir la liste des rangs de chaque groupe. Voici la construction du groupe TP :

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

Au moment de la communication, les couches modèle ne touchent pas directement ProcessGroup : elles appellent get_tp_group().all_reduce(tensor). À l'intérieur, all_reduce choisit entre custom op et appel direct à device_communicator.all_reduce selon use_custom_op_call :

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

Le forward de ColumnParallelLinear appelle all_reduce après avoir calculé la part locale pour fusionner les fragments ; RowParallelLinear fait l'inverse : all-gather d'abord, calcul local, puis reduce_scatter. Toutes ces ops sont encapsulées au niveau de GroupCoordinator.

Limites et échecs

  • rank_in_group obligatoire pour le découpage : les couches modèle doivent découper avec rank_in_group, pas le rank global, sinon en multi-nœud elles tomberaient sur le mauvais fragment(parallel_state.py:380).
  • Initialisation TP répétée → assertion : initialize_model_parallel vérifie assert _TP is None pour chaque groupe(parallel_state.py:1798) ; un second appel lève une AssertionError, il faut d'abord destroy_model_parallel.
  • DCP ne peut pas dépasser TP : dcp_size <= tp_size, car DCP réutilise les GPU du groupe TP(parallel_state.py:1816-1820).
  • Mono-GPU court-circuited : all_reduce retourne l'entrée quand world_size == 1(parallel_state.py:657-658) sans processus NCCL ; en revanche, si device_communicator is None hors cas bypass, lève un ValueError.
  • DeviceCommunicator manquant → erreur : _all_reduce_out_place vérifie device_communicator is None et lève(parallel_state.py:665-668), pour éviter un appel avant initialisation du backend.
  • _replace_active_groups doit être collectif : parallel_state.py:1341-1362 exige que tous les rangs appellent en même temps, sinon deadlock ; en EP élastique, le supervisor orchestre l'appel.

Résumé

parallel_state est la fondation de l'inférence multi-GPU de vLLM : GroupCoordinator encapsule le ProcessGroup de PyTorch en regroupant cpu_group/device_group, custom ops, device_communicator et broadcaster par file de messages ; les couches modèle récupèrent le groupe via get_tp_group() / get_pp_group() puis appellent all_reduce et cie. L'utilisation de ces communications par les couches linéaires TP est détaillée dans /model-loading/tp-layers ; le démarrage des groupes côté Worker dans /worker/worker ; l'assemblage de ces workers par le cœur du moteur dans /engine/engine-core.

Voir la documentation officielle : Documentation vLLM · README