Capas lineales con tensor parallelism: ColumnParallelLinear / QKVParallelLinear / RowParallelLinear
Responsabilidades
El tensor parallelism (TP) de vLLM parte un único nn.Linear en varias GPUs, cada una calcula una parte, y en el forward / backward el resultado completo se reensambla con all-gather / all-reduce. Esta capa de abstracción reside en tres clases centrales de vllm/model_executor/layers/linear.py: ColumnParallelLinear parte a lo largo de la dimensión de salida, RowParallelLinear a lo largo de la dimensión de entrada, y QKVParallelLinear gestiona los tres segmentos Q/K/V con distinto número de heads en attention (sobre todo en GQA / MQA). Las tres heredan de LinearBase y delegan la llamada GEMM real a quant_method.apply (linear.py:551-569), de modo que la misma lógica de paralelismo sirve para backends sin cuantizar, fp8, AWQ, GPTQ, etc.
Al instanciar cada capa lineal TP se calculan el tp_rank / tp_size de este rank, se divide output_size entre tp_size para obtener output_size_per_partition (linear.py:437-442) y se escribe esa información de partición en los atributos output_dim / input_dim de los parámetros; el weight_loader usa luego ese atributo para decidir qué segmento cortar. QKVParallelLinear añade encima la replicación de heads KV para GQA: cuando tp_size >= total_num_kv_heads, cada rank mantiene una copia completa de los heads KV (num_kv_head_replicas = tp_size / total_num_kv_heads) (linear.py:994-1001).
Motivación de diseño
¿Por qué introducir encima de LinearBase dos mecanismos, weight_loader / weight_loader_v2?
- Corte distinto:
ColumnParallelLinearcorta por la dimensión de salida (cada rank toma un segmento de[A_1, A_2, ...]),RowParallelLinearpor la dimensión de entrada (A = [A_1; A_2; ...],X = [X_1, X_2, ...]), por lo que losweight_loaderde cada uno hacennarrowsegúnoutput_dimoinput_dim(linear.py:530-533)(linear.py:1650-1653). - v2 va por objetos de parámetros: el nuevo
weight_loader_v2envuelve el tensor en unBasevLLMParametery llama directamente aparam.load_column_parallel_weight/param.load_row_parallel_weight, encapsulando "qué segmento cortar" dentro del propio parámetro (linear.py:543-549)(linear.py:1663-1670); los métodos de cuantización pueden elegir v2 (WEIGHT_LOADER_V2_SUPPORTED) o la ruta vieja. - Fusión de los tres segmentos Q/K/V:
QKVParallelLinearconcatena los pesos de Q/K/V en la dimensión de salida formando una matriz grande;output_sizes = [q, k, v]calcula cada segmento condivide(..., tp_size)por separado (linear.py:1003-1008);MergedColumnParallelLineares la versión generalizada, capaz de fusionar cualquier conjunto de capas column-parallel. - Replicación de KV en GQA: en
QKVParallelLinear.__init__se ejecutaif 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) para que cada rank pueda calcular un head KV completo y el forward posterior no necesite sincronización extra. - Metadatos de shard: el
weight_loader_v2deColumnParallelLineartambién reconoce elshard_id(en QKV es"q" / "k" / "v") y según el shard id decide a qué offset del param escribir (linear.py:1023-1047). - bias solo en rank 0: en
RowParallelLinear.forward,bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias(linear.py:1685-1688); sumar el bias después del reduce duplicaría, así que solo se añade en rank 0. - TP desactivable: con
disable_tp=Truese fuerzatp_rank=0/tp_size=1y la capa degenera en unLinearnormal, útil para ciertos experts no paralelos de MoE (linear.py:438-439).
Archivos clave
LinearMethodBase:141-180— clase abstracta base de los métodos de cuantización; define las tres interfacescreate_weights/apply/process_weights_after_loading.LinearBase:231-288— padre de todas las capas lineales;update_param_tp_statusescribetp_rank/tp_sizeen cadaBasevLLMParameter.ColumnParallelLinear.__init__:423-491— parte por la dimensión de salida:output_size_per_partition = divide(output_size, tp_size).ColumnParallelLinear weight_loader:520-549— v1 hace narrow segúnoutput_dim, v2 delega enparam.load_column_parallel_weight.ColumnParallelLinear.forward:551-569—quant_method.apply→tensor_model_parallel_all_gatheropcional.QKVParallelLinear.__init__:942-1021— gestiona la replicación de heads KV en GQA y concatena los tresoutput_sizesde Q/K/V en unoutput_size.QKVParallelLinear shard mapping:1023-1047—_get_shard_offset_mapping/_get_shard_size_mappingle dicen al weight_loader a qué segmento del param escribir.QKVParallelLinear fused checkpoint:1049-1097— para QKV ya fusionado en disco (tipo Phi-3), lo parte automáticamente y llama aweight_loader_v2segmento por segmento.RowParallelLinear.__init__:1572-1626— parte por la dimensión de entrada:input_size_per_partition = divide(input_size, tp_size).RowParallelLinear.forward:1672-1698— si la entrada no viene partida, primero split; tras el GEMM,tensor_model_parallel_all_reduce.MergedColumnParallelLinear:580-636— versión generalizada de la fusión de varias capas column-parallel (la base de QKV);output_sizesdefine cada segmento.
Flujo de datos
Lo siguiente es la lógica central de ColumnParallelLinear.weight_loader: tras recibir el tensor completo desde el disco, lo corta por el offset tp_rank * shard_size para quedarse con el segmento de este rank y lo copia con copy_ al parámetro:
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 es la imagen espejo: cambia la dimensión del narrow de output_dim a input_dim (linear.py:1650-1653). DefaultModelLoader los invoca indirectamente al cargar los pesos vía model.load_weights —cada implementación de modelo tiene algo parecido a loaded_weights = param.weight_loader(param, loaded_weight).
En el forward, ColumnParallelLinear hace all-gather (con gather_output=True concatena el output_parallel de cada rank en un output completo) y RowParallelLinear hace all-reduce (con reduce_results=True suma los resultados del GEMM fragmentado); cuando se combinan ambos, ahorran una sincronización: QKVParallelLinear (column) → attention → o_proj (row) es la combinación clásica column-then-row.
Límites y fallos
- shard_id inválido en QKV:
validate_shard_idrestringe el shard id de QKV a"q" / "k" / "v"oNone; cualquier otro valor lanzaValueErrordirectamente (linear.py:1023-1029). - RowParallelLinear con
reduce_results=False+ bias: no hacer reduce y a la vez añadir bias dentro del GEMM duplicaría el bias en cada rank; el código lanzaValueError(linear.py:1622-1626). - Forma de bloque FP8 incompatible:
_maybe_allow_fp8_block_shape_mismatchrelajaallow_fp8_block_shape_mismatchcuandoweight_block_sizeno divide algúnoutput_partition_sizes(linear.py:493-516); en caso contrario el kernel rechaza la forma. - bitsandbytes salta el narrow: con
use_bitsandbytes_4bit, el checkpoint ya viene partido por rank, así que se hacecopy_directo sin narrow (linear.py:524-527). - scales como escalar de dimensión cero: checkpoints como AutoFP8 guardan las scales como tensores de dimensión cero;
weight_loaderlas reemplaza conreshape(1)como fallback (linear.py:537-538). - capas row con
input_is_parallel=False: si la entrada no está partida,RowParallelLinear.forwardhacesplit_tensor_along_last_dimpara tomar su segmento (linear.py:1676-1682).
Resumen
Las tres clases lineales paralelas forman el "lenguaje de partición de pesos" del TP de vLLM; el cálculo concreto lo gestiona quant_method. El contacto con carga de capas cuantizadas es: ColumnParallelLinear.__init__ llama a self.quant_method.create_weights(...) para insertar el weight_loader (v1 o v2) en cada parámetro; el contacto con DefaultModelLoader es que model.load_weights invoca weight_loader para cada tensor y corta el peso completo del disco en el segmento correspondiente a este rank.