Couches linéaires en parallélisme de tenseurs : ColumnParallelLinear / QKVParallelLinear / RowParallelLinear
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 :
ColumnParallelLinearcoupe selon la dimension de sortie (chaque rank reçoit un segment de[A_1, A_2, ...]),RowParallelLinearselon la dimension d'entrée (A = [A_1; A_2; ...],X = [X_1, X_2, ...]). Les deuxweight_loaderfont donc unnarrowselonoutput_dimouinput_dim(linear.py:530-533)(linear.py:1650-1653). - v2 s'appuie sur l'objet paramètre : le nouveau
weight_loader_v2enveloppe le tenseur dans unBasevLLMParameteret appelle directementparam.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 :
QKVParallelLinearconcatè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 sondivide(..., tp_size)(linear.py:1003-1008).MergedColumnParallelLinearen 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_v2deColumnParallelLineardoit aussi reconnaître leshard_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 forcetp_rank=0/tp_size=1et la couche redevient unLinearordinaire — utile pour certains experts non parallèles des MoE(linear.py:438-439).
Fichiers clés
LinearMethodBase:141-180— classe de base abstraite des méthodes de quantification, définit les trois interfacescreate_weights/apply/process_weights_after_loading.LinearBase:231-288— parent de toutes les couches linéaires ;update_param_tp_statusécrittp_rank/tp_sizesur chaqueBasevLLMParameter.ColumnParallelLinear.__init__:423-491— découpe selon la dimension de sortie,output_size_per_partition = divide(output_size, tp_size).ColumnParallelLinear weight_loader:520-549— v1 fait unnarrowsuroutput_dim, v2 délègue àparam.load_column_parallel_weight.ColumnParallelLinear.forward:551-569—quant_method.apply→ optionneltensor_model_parallel_all_gather.QKVParallelLinear.__init__:942-1021— gère la réplique des KV heads en GQA, concatène lesoutput_sizesdes trois segments Q/K/V enoutput_size.QKVParallelLinear shard mapping:1023-1047—_get_shard_offset_mapping/_get_shard_size_mapping, indiquent au weight_loader à quel segment du paramètre écrire.QKVParallelLinear fused checkpoint:1049-1097— pour les checkpoints QKV déjà fusionnés sur disque (Phi-3), découpe automatiquement puis appelleweight_loader_v2par segment.RowParallelLinear.__init__:1572-1626— découpe selon la dimension d'entrée,input_size_per_partition = divide(input_size, tp_size).RowParallelLinear.forward:1672-1698— si l'entrée n'est pas déjà parallélisée, on la split d'abord ; après le GEMM,tensor_model_parallel_all_reduce.MergedColumnParallelLinear:580-636— fusion de plusieurs couches parallèles-colonnes (version générique de QKV),output_sizesfixe chaque segment.
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 :
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_idrestreint les shard id QKV à"q" / "k" / "v"ouNone, toute autre valeur lève unValueError(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_mismatchactiveallow_fp8_block_shape_mismatchquandweight_block_sizene divise pas unoutput_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 directementcopy_sans narrow(linear.py:524-527). - scales en tenseur zéro-dimensionnel : les checkpoints AutoFP8 stockent les scales comme tenseurs 0D ;
weight_loaders'en sort avecreshape(1)(linear.py:537-538). - Couche row avec input_is_parallel=False : quand l'entrée n'est pas découpée,
RowParallelLinear.forwardlui-même faitsplit_tensor_along_last_dimpour 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