Skip to content

Couches linéaires en parallélisme de tenseurs : ColumnParallelLinear / QKVParallelLinear / RowParallelLinear

源码版本v0.25.1

Responsabilités

Le parallélisme de tenseurs (tensor parallelism, TP) de vLLM découpe un nn.Linear en plusieurs GPU, chacun calculant une partie, puis réassemble le résultat complet via all-gather / all-reduce en forward et en backward. Cette abstraction se matérialise dans vllm/model_executor/layers/linear.py à travers trois classes centrales : ColumnParallelLinear (découpe selon la dimension de sortie), RowParallelLinear (découpe selon la dimension d'entrée) et QKVParallelLinear (gère les trois segments Q/K/V de l'attention avec des nombres de heads différents, notamment pour GQA / MQA). Toutes héritent de LinearBase et délèguent le véritable GEMM à quant_method.apply(linear.py:551-569) : la même logique de parallélisation s'applique donc à tous les backends, non quantifié, fp8, AWQ, GPTQ, etc.

À l'instanciation, chaque couche linéaire TP calcule le tp_rank / tp_size du rank courant, divise output_size par tp_size pour obtenir output_size_per_partition(linear.py:437-442), et écrit cette information de partitionnement dans les attributs output_dim / input_dim des paramètres. Le weight_loader s'appuie ensuite sur cet attribut pour choisir quel segment extraire. QKVParallelLinear ajoute à cela la réplique des KV heads en GQA : quand tp_size >= total_num_kv_heads, chaque rank détient une copie complète des KV heads (num_kv_head_replicas = tp_size / total_num_kv_heads)(linear.py:994-1001).

Motivation de conception

Pourquoi introduire deux mécanismes weight_loader / weight_loader_v2 au-dessus de LinearBase ?

  • Découpage différent : ColumnParallelLinear coupe selon la dimension de sortie (chaque rank reçoit un segment de [A_1, A_2, ...]), RowParallelLinear selon la dimension d'entrée (A = [A_1; A_2; ...], X = [X_1, X_2, ...]). Les deux weight_loader font donc un narrow selon output_dim ou input_dim(linear.py:530-533)(linear.py:1650-1653).
  • v2 s'appuie sur l'objet paramètre : le nouveau weight_loader_v2 enveloppe le tenseur dans un BasevLLMParameter et appelle directement param.load_column_parallel_weight / param.load_row_parallel_weight, en encapsulant « quel segment extraire » dans l'objet paramètre(linear.py:543-549)(linear.py:1663-1670). La méthode de quantification peut opter pour v2 (WEIGHT_LOADER_V2_SUPPORTED) ou rester sur l'ancien chemin.
  • Fusion des trois segments Q/K/V : QKVParallelLinear concatène les poids de Q/K/V en une seule matrice sur la dimension de sortie, output_sizes = [q, k, v] où chaque segment calcule son divide(..., tp_size)(linear.py:1003-1008). MergedColumnParallelLinear en est la version générique, capable de fusionner n'importe quel ensemble de couches parallèles-colonnes.
  • Réplique KV en GQA : dans QKVParallelLinear.__init__, if tp_size >= total_num_kv_heads: self.num_kv_heads = 1; self.num_kv_head_replicas = divide(tp_size, total_num_kv_heads)(linear.py:996-1001) garantit que chaque rank calcule un KV head complet, sans synchronisation supplémentaire au forward.
  • Métadonnées de shard : le weight_loader_v2 de ColumnParallelLinear doit aussi reconnaître le shard_id ("q" / "k" / "v" pour QKV) et décider selon celui-ci à quel offset du paramètre écrire(linear.py:1023-1047).
  • bias uniquement sur rank 0 : dans RowParallelLinear.forward, bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias(linear.py:1685-1688) : ajouter le bias après le reduce le doublerait, donc on ne l'ajoute que sur rank 0.
  • TP désactivable : avec disable_tp=True, on force tp_rank=0 / tp_size=1 et la couche redevient un Linear ordinaire — utile pour certains experts non parallèles des MoE(linear.py:438-439).

Fichiers clés

Flux de données

Voici la logique centrale de ColumnParallelLinear.weight_loader — une fois le tenseur complet lu depuis le disque, on découpe selon l'offset tp_rank * shard_size pour obtenir le segment du rank courant, puis copy_ dans le paramètre :

python
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
    output_dim = getattr(param, "output_dim", None)

    is_sharded_weight = getattr(param, "is_sharded_weight", False)
    use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
    # bitsandbytes loads the weights of the specific portion
    # no need to narrow
    is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit

    param_data = param.data
    if output_dim is not None and not is_sharded_weight:
        shard_size = param_data.shape[output_dim]
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)

    # Special case for loading scales off disk, which often do not
    # have a shape (such as in the case of AutoFP8).
    if len(loaded_weight.shape) == 0:
        loaded_weight = loaded_weight.reshape(1)

    assert param_data.shape == loaded_weight.shape
    param_data.copy_(loaded_weight)

(linear.py:520-541) RowParallelLinear.weight_loader en est le miroir, simplement la dimension du narrow passe de output_dim à input_dim(linear.py:1650-1653). Ces deux fonctions sont appelées indirectement par DefaultModelLoader via model.load_weights — chaque implémentation de modèle contient quelque chose comme loaded_weights = param.weight_loader(param, loaded_weight).

Au forward, ColumnParallelLinear fait un all-gather (avec gather_output=True, les output_parallel de chaque rank sont concaténés en un output complet), tandis que RowParallelLinear fait un all-reduce (avec reduce_results=True, les résultats des GEMM par bloc sont sommés). Utilisées ensemble, elles permettent d'éliminer une synchronisation : QKVParallelLinear (column) → attention → o_proj (row) est la combinaison classique column-then-row.

Limites et échecs

  • shard_id QKV invalide : validate_shard_id restreint les shard id QKV à "q" / "k" / "v" ou None, toute autre valeur lève un ValueError(linear.py:1023-1029).
  • RowParallelLinear reduce_results=False + bias : ne pas réduire tout en ajoutant le bias dans le GEMM ferait que chaque rank l'ajoute, d'où un résultat erroné ; le code lève un ValueError(linear.py:1622-1626).
  • Forme de bloc FP8 incompatible : _maybe_allow_fp8_block_shape_mismatch active allow_fp8_block_shape_mismatch quand weight_block_size ne divise pas un output_partition_sizes(linear.py:493-516), sinon le kernel refusera.
  • bitsandbytes ignore le narrow : avec use_bitsandbytes_4bit, le checkpoint est déjà partitionné par rank, on fait directement copy_ sans narrow(linear.py:524-527).
  • scales en tenseur zéro-dimensionnel : les checkpoints AutoFP8 stockent les scales comme tenseurs 0D ; weight_loader s'en sort avec reshape(1)(linear.py:537-538).
  • Couche row avec input_is_parallel=False : quand l'entrée n'est pas découpée, RowParallelLinear.forward lui-même fait split_tensor_along_last_dim pour prendre la part du rank courant(linear.py:1676-1682).

Résumé

Les trois classes de couches linéaires parallèles constituent le « langage de découpe de poids » du TP de vLLM ; la manière de calculer est laissée à quant_method. L'articulation avec chargement des couches quantifiées se fait dans ColumnParallelLinear.__init__, qui appelle self.quant_method.create_weights(...) et injecte le weight_loader (v1 ou v2) dans chaque paramètre. L'articulation avec DefaultModelLoader se fait dans model.load_weights, qui appelle weight_loader tenseur par tenseur pour découper le poids complet du disque en la part du rank courant.

Voir la documentation officielle : Documentation vLLM · README